From a1dc6932182157224d3d8ca348273ac155ebd01f Mon Sep 17 00:00:00 2001 From: Abhinav Rastogi Date: Wed, 7 Oct 2026 02:34:56 +0530 Subject: [PATCH 1/4] fix: avoid failover on client-local errors --- src/typesense/async_/api_call.py | 16 ++++++++++++++++ src/typesense/sync/api_call.py | 16 ++++++++++++++++ tests/api_call_test.py | 19 +++++++++++++++++++ 3 files changed, 51 insertions(+) diff --git a/src/typesense/async_/api_call.py b/src/typesense/async_/api_call.py index be1a83d..a85f6a7 100644 --- a/src/typesense/async_/api_call.py +++ b/src/typesense/async_/api_call.py @@ -135,6 +135,20 @@ class SessionFunctionKwargs(typing.Generic[TParams, TBody], typing.TypedDict): ServiceUnavailable, ) +_CLIENT_ERRORS: typing.Final[ + typing.Tuple[ + typing.Type[httpx.PoolTimeout], + typing.Type[httpx.LocalProtocolError], + typing.Type[httpx.DecodingError], + typing.Type[httpx.TooManyRedirects], + ] +] = ( + httpx.PoolTimeout, + httpx.LocalProtocolError, + httpx.DecodingError, + httpx.TooManyRedirects, +) + class AsyncApiCall: """ @@ -478,6 +492,8 @@ async def _execute_request( as_json, **request_kwargs, ) + except _CLIENT_ERRORS: + raise except _SERVER_ERRORS as server_error: self.node_manager.set_node_health(node, is_healthy=False) if num_retries < self.config.num_retries: diff --git a/src/typesense/sync/api_call.py b/src/typesense/sync/api_call.py index 402a0dc..1290774 100644 --- a/src/typesense/sync/api_call.py +++ b/src/typesense/sync/api_call.py @@ -135,6 +135,20 @@ class SessionFunctionKwargs(typing.Generic[TParams, TBody], typing.TypedDict): ServiceUnavailable, ) +_CLIENT_ERRORS: typing.Final[ + typing.Tuple[ + typing.Type[httpx.PoolTimeout], + typing.Type[httpx.LocalProtocolError], + typing.Type[httpx.DecodingError], + typing.Type[httpx.TooManyRedirects], + ] +] = ( + httpx.PoolTimeout, + httpx.LocalProtocolError, + httpx.DecodingError, + httpx.TooManyRedirects, +) + class ApiCall: """ @@ -478,6 +492,8 @@ def _execute_request( as_json, **request_kwargs, ) + except _CLIENT_ERRORS: + raise except _SERVER_ERRORS as server_error: self.node_manager.set_node_health(node, is_healthy=False) if num_retries < self.config.num_retries: diff --git a/tests/api_call_test.py b/tests/api_call_test.py index b7c4888..280bb33 100644 --- a/tests/api_call_test.py +++ b/tests/api_call_test.py @@ -461,6 +461,25 @@ def test_selects_next_available_node_on_timeout( assert len(respx.calls) == 3 +def test_client_errors_do_not_mark_nodes_unhealthy( + fake_api_call: ApiCall, + mocker: MockerFixture, +) -> None: + """Pool exhaustion is local to the client and must not trigger failover.""" + node = fake_api_call.node_manager.get_node() + make_request = mocker.patch.object( + fake_api_call.request_handler, + "make_request", + side_effect=httpx.PoolTimeout("No connection available"), + ) + + with pytest.raises(httpx.PoolTimeout): + fake_api_call.get("/test", as_json=True, entity_type=typing.Dict[str, str]) + + assert node.healthy is True + make_request.assert_called_once() + + def test_get_node_no_healthy_nodes( fake_api_call: ApiCall, mocker: MockFixture, From a9be046056efdf85fbde441799934d5f572a7c37 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Wed, 7 Oct 2026 13:08:41 +0300 Subject: [PATCH 2/4] test: cover every client-local error on both clients (#143) --- tests/api_call_test.py | 58 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 58 insertions(+) diff --git a/tests/api_call_test.py b/tests/api_call_test.py index 280bb33..ea60b5b 100644 --- a/tests/api_call_test.py +++ b/tests/api_call_test.py @@ -684,3 +684,61 @@ async def test_async_sleeps_retry_interval_between_retries( assert sleep_call == mocker.call( fake_async_api_call.config.retry_interval_seconds, ) + + +@pytest.mark.parametrize( + "client_side_error", + [ + httpx.PoolTimeout("Pool timeout"), + httpx.LocalProtocolError("Local protocol error"), + httpx.DecodingError("Decoding error"), + httpx.TooManyRedirects("Too many redirects"), + ], +) +def test_client_side_error_does_not_mark_node_unhealthy( + fake_api_call: ApiCall, + client_side_error: httpx.HTTPError, +) -> None: + """Test that client-side httpx errors propagate without failing over.""" + with respx.mock: + respx.get("http://nearest:8108/").mock(side_effect=client_side_error) + node0_route = respx.get("http://node0:8108/").mock( + return_value=httpx.Response(200, json={"key": "value"}), + ) + + with pytest.raises(type(client_side_error)): + fake_api_call.get("/", entity_type=typing.Dict[str, str]) + + assert len(respx.calls) == 1 + assert not node0_route.called + + assert fake_api_call.config.nearest_node.healthy is True + + +@pytest.mark.parametrize( + "client_side_error", + [ + httpx.PoolTimeout("Pool timeout"), + httpx.LocalProtocolError("Local protocol error"), + httpx.DecodingError("Decoding error"), + httpx.TooManyRedirects("Too many redirects"), + ], +) +async def test_async_client_side_error_does_not_mark_node_unhealthy( + fake_async_api_call: AsyncApiCall, + client_side_error: httpx.HTTPError, +) -> None: + """Test that client-side httpx errors propagate without failing over (async).""" + with respx.mock: + respx.get("http://nearest:8108/").mock(side_effect=client_side_error) + node0_route = respx.get("http://node0:8108/").mock( + return_value=httpx.Response(200, json={"key": "value"}), + ) + + with pytest.raises(type(client_side_error)): + await fake_async_api_call.get("/", entity_type=typing.Dict[str, str]) + + assert len(respx.calls) == 1 + assert not node0_route.called + + assert fake_async_api_call.config.nearest_node.healthy is True From 3d0d70c03cbfdd48bde619229832eb44bd188b4a Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Tue, 6 Oct 2026 13:38:13 +0300 Subject: [PATCH 3/4] fix: mark the node that answered as healthy (#144) --- src/typesense/async_/api_call.py | 9 ++-- src/typesense/sync/api_call.py | 9 ++-- tests/api_call_test.py | 70 ++++++++++++++++++++++++++++++++ 3 files changed, 78 insertions(+), 10 deletions(-) diff --git a/src/typesense/async_/api_call.py b/src/typesense/async_/api_call.py index a85f6a7..8ad68ec 100644 --- a/src/typesense/async_/api_call.py +++ b/src/typesense/async_/api_call.py @@ -487,6 +487,7 @@ async def _execute_request( try: return await self._make_request_and_process_response( method, + node, url, entity_type, as_json, @@ -511,12 +512,13 @@ async def _execute_request( async def _make_request_and_process_response( self, method: str, + node: Node, url: str, entity_type: typing.Type[TEntityDict], as_json: bool, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> typing.Union[TEntityDict, str]: - """Make the async API request and process the response.""" + """Make the async API request to `node` and process the response.""" request_response = await self.request_handler.make_request( method=method, url=url, @@ -525,10 +527,7 @@ async def _make_request_and_process_response( client=self._client, **kwargs, ) - self.node_manager.set_node_health( - self.node_manager.get_node(), - is_healthy=True, - ) + self.node_manager.set_node_health(node, is_healthy=True) return ( typing.cast(TEntityDict, request_response) if as_json diff --git a/src/typesense/sync/api_call.py b/src/typesense/sync/api_call.py index 1290774..f65d320 100644 --- a/src/typesense/sync/api_call.py +++ b/src/typesense/sync/api_call.py @@ -487,6 +487,7 @@ def _execute_request( try: return self._make_request_and_process_response( method, + node, url, entity_type, as_json, @@ -511,12 +512,13 @@ def _execute_request( def _make_request_and_process_response( self, method: str, + node: Node, url: str, entity_type: typing.Type[TEntityDict], as_json: bool, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> typing.Union[TEntityDict, str]: - """Make the async API request and process the response.""" + """Make the async API request to `node` and process the response.""" request_response = self.request_handler.make_request( method=method, url=url, @@ -525,10 +527,7 @@ def _make_request_and_process_response( client=self._client, **kwargs, ) - self.node_manager.set_node_health( - self.node_manager.get_node(), - is_healthy=True, - ) + self.node_manager.set_node_health(node, is_healthy=True) return ( typing.cast(TEntityDict, request_response) if as_json diff --git a/tests/api_call_test.py b/tests/api_call_test.py index ea60b5b..fb4857b 100644 --- a/tests/api_call_test.py +++ b/tests/api_call_test.py @@ -742,3 +742,73 @@ async def test_async_client_side_error_does_not_mark_node_unhealthy( assert not node0_route.called assert fake_async_api_call.config.nearest_node.healthy is True + + +def test_round_robin_visits_each_node_in_turn(fake_api_call: ApiCall) -> None: + """Test that successful requests advance the round-robin by one node each.""" + fake_api_call.config.nearest_node = None + + with respx.mock: + for host in ("node0", "node1", "node2"): + respx.get(f"http://{host}:8108/").mock( + return_value=httpx.Response(200, json={"key": "value"}), + ) + + for _ in range(6): + fake_api_call.get("/", entity_type=typing.Dict[str, str]) + + assert [str(call.request.url) for call in respx.calls] == [ + "http://node0:8108/", + "http://node1:8108/", + "http://node2:8108/", + "http://node0:8108/", + "http://node1:8108/", + "http://node2:8108/", + ] + + +async def test_async_round_robin_visits_each_node_in_turn( + fake_async_api_call: AsyncApiCall, +) -> None: + """Test that successful requests advance the round-robin by one node each (async).""" + fake_async_api_call.config.nearest_node = None + + with respx.mock: + for host in ("node0", "node1", "node2"): + respx.get(f"http://{host}:8108/").mock( + return_value=httpx.Response(200, json={"key": "value"}), + ) + + for _ in range(6): + await fake_async_api_call.get("/", entity_type=typing.Dict[str, str]) + + assert [str(call.request.url) for call in respx.calls] == [ + "http://node0:8108/", + "http://node1:8108/", + "http://node2:8108/", + "http://node0:8108/", + "http://node1:8108/", + "http://node2:8108/", + ] + + +def test_success_marks_only_the_answering_node_healthy( + fake_api_call: ApiCall, +) -> None: + """Test that a success refreshes the node that answered and no other.""" + fake_api_call.config.nearest_node = None + answering_node, unhealthy_node, _ = fake_api_call.node_manager.nodes + answering_node.last_access_ts = 0 + unhealthy_node.healthy = False + unhealthy_node.last_access_ts = int(time.time()) + + with respx.mock: + respx.get("http://node0:8108/").mock( + return_value=httpx.Response(200, json={"key": "value"}), + ) + + fake_api_call.get("/", entity_type=typing.Dict[str, str]) + + assert answering_node.healthy is True + assert answering_node.last_access_ts > 0 + assert unhealthy_node.healthy is False From 6aae5005d66b7a6c64bdd771fffaa9b0025cac38 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Tue, 6 Oct 2026 13:58:47 +0300 Subject: [PATCH 4/4] feat: configure the connection pool and cap requests in flight (#146) --- src/typesense/async_/api_call.py | 30 ++++-- src/typesense/concurrency_limit.py | 89 ++++++++++++++++ src/typesense/configuration.py | 63 ++++++++++++ src/typesense/sync/api_call.py | 30 ++++-- tests/api_call_test.py | 131 ++++++++++++++++++++++++ tests/configuration_test.py | 43 ++++++++ tests/configuration_validations_test.py | 37 +++++++ utils/run-unasync.py | 2 + 8 files changed, 407 insertions(+), 18 deletions(-) create mode 100644 src/typesense/concurrency_limit.py diff --git a/src/typesense/async_/api_call.py b/src/typesense/async_/api_call.py index 8ad68ec..916b9bb 100644 --- a/src/typesense/async_/api_call.py +++ b/src/typesense/async_/api_call.py @@ -37,6 +37,7 @@ import httpx +from typesense.concurrency_limit import AsyncConcurrencyLimit from typesense.configuration import Configuration, Node from typesense.exceptions import ( HTTPStatus0Error, @@ -174,7 +175,17 @@ def __init__(self, config: Configuration): self.node_manager = NodeManager(config) self.request_handler = RequestHandler(config) self._client = httpx.AsyncClient( - timeout=config.connection_timeout_seconds, + timeout=httpx.Timeout( + config.connection_timeout_seconds, + pool=config.pool_timeout_seconds, + ), + limits=httpx.Limits( + max_connections=config.max_connections, + max_keepalive_connections=config.max_keepalive_connections, + ), + ) + self._concurrency_limit = AsyncConcurrencyLimit( + config.max_concurrent_requests, ) async def __aenter__(self) -> "AsyncApiCall": @@ -519,14 +530,15 @@ async def _make_request_and_process_response( **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> typing.Union[TEntityDict, str]: """Make the async API request to `node` and process the response.""" - request_response = await self.request_handler.make_request( - method=method, - url=url, - as_json=as_json, - entity_type=entity_type, - client=self._client, - **kwargs, - ) + async with self._concurrency_limit: + request_response = await self.request_handler.make_request( + method=method, + url=url, + as_json=as_json, + entity_type=entity_type, + client=self._client, + **kwargs, + ) self.node_manager.set_node_health(node, is_healthy=True) return ( typing.cast(TEntityDict, request_response) diff --git a/src/typesense/concurrency_limit.py b/src/typesense/concurrency_limit.py new file mode 100644 index 0000000..34ebdc5 --- /dev/null +++ b/src/typesense/concurrency_limit.py @@ -0,0 +1,89 @@ +""" +Optional caps on the number of requests a client sends at once. + +``AsyncConcurrencyLimit`` is used by the async client and ``ConcurrencyLimit`` by the +sync client (``utils/run-unasync.py`` maps one name to the other). Both are no-ops +when ``max_concurrent_requests`` is ``None``. + +Keeping the cap below the httpx pool's ``max_connections`` means requests queue here +instead of in the pool, so a burst of slow requests cannot exhaust the pool and +raise ``httpx.PoolTimeout``. +""" + +import asyncio +import sys +import threading +from types import TracebackType + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + + +class AsyncConcurrencyLimit: + """Async context manager that holds a slot for the duration of a request.""" + + def __init__(self, max_concurrent_requests: typing.Optional[int]) -> None: + """ + Initialize the limit. + + Args: + max_concurrent_requests (Optional[int]): The maximum number of requests + in flight at once, or ``None`` for no limit. + """ + self._max_concurrent_requests = max_concurrent_requests + # Created on first use, inside the running event loop. On Python < 3.10 a + # semaphore binds to the loop that is current when it is constructed. + self._semaphore: typing.Optional[asyncio.Semaphore] = None + + async def __aenter__(self) -> None: + """Wait for a free slot.""" + if self._max_concurrent_requests is None: + return + if self._semaphore is None: + self._semaphore = asyncio.Semaphore(self._max_concurrent_requests) + await self._semaphore.acquire() + + async def __aexit__( + self, + exc_type: typing.Optional[typing.Type[BaseException]], + exc_val: typing.Optional[BaseException], + exc_tb: typing.Optional[TracebackType], + ) -> None: + """Release the slot.""" + if self._semaphore is not None: + self._semaphore.release() + + +class ConcurrencyLimit: + """Context manager that holds a slot for the duration of a request.""" + + def __init__(self, max_concurrent_requests: typing.Optional[int]) -> None: + """ + Initialize the limit. + + Args: + max_concurrent_requests (Optional[int]): The maximum number of requests + in flight at once, or ``None`` for no limit. + """ + self._semaphore: typing.Optional[threading.Semaphore] = ( + None + if max_concurrent_requests is None + else threading.Semaphore(max_concurrent_requests) + ) + + def __enter__(self) -> None: + """Wait for a free slot.""" + if self._semaphore is not None: + self._semaphore.acquire() + + def __exit__( + self, + exc_type: typing.Optional[typing.Type[BaseException]], + exc_val: typing.Optional[BaseException], + exc_tb: typing.Optional[TracebackType], + ) -> None: + """Release the slot.""" + if self._semaphore is not None: + self._semaphore.release() diff --git a/src/typesense/configuration.py b/src/typesense/configuration.py index aaa741e..4f9144a 100644 --- a/src/typesense/configuration.py +++ b/src/typesense/configuration.py @@ -82,6 +82,23 @@ class ConfigDict(typing.TypedDict): connection_timeout_seconds (float): The connection timeout in seconds. suppress_deprecation_warnings (bool): Whether to suppress deprecation warnings. + + pool_timeout_seconds (float): How long a request waits for a free connection + in the pool before raising ``httpx.PoolTimeout``. Defaults to + ``connection_timeout_seconds``. Setting it lower than + ``connection_timeout_seconds`` makes the httpcore connection leak + (encode/httpcore#1093) more likely under load. + + max_connections (int): The maximum number of connections in the pool. + Defaults to 100. + + max_keepalive_connections (int): The maximum number of idle connections + kept alive in the pool. Defaults to 20. + + max_concurrent_requests (int): The maximum number of requests in flight at + once; further requests wait for a slot. Keep it below + ``max_connections`` so a burst of slow requests cannot exhaust the pool. + Defaults to no limit. """ nodes: typing.List[typing.Union[str, NodeConfigDict]] @@ -100,6 +117,10 @@ class ConfigDict(typing.TypedDict): ] # deprecated connection_timeout_seconds: typing.NotRequired[float] suppress_deprecation_warnings: typing.NotRequired[bool] + pool_timeout_seconds: typing.NotRequired[float] + max_connections: typing.NotRequired[int] + max_keepalive_connections: typing.NotRequired[int] + max_concurrent_requests: typing.NotRequired[int] class Node: @@ -188,6 +209,10 @@ class Configuration: retry_interval_seconds (float): The interval in seconds between retries. healthcheck_interval_seconds (int): The interval in seconds between health checks. verify (bool): Whether to verify the SSL certificate. + pool_timeout_seconds (float): How long to wait for a free pooled connection. + max_connections (int): The maximum number of connections in the pool. + max_keepalive_connections (int): The maximum number of idle pooled connections. + max_concurrent_requests (int | None): The maximum number of requests in flight. """ def __init__( @@ -232,6 +257,18 @@ def __init__( self.suppress_deprecation_warnings = config_dict.get( "suppress_deprecation_warnings", False ) + self.pool_timeout_seconds = config_dict.get( + "pool_timeout_seconds", + self.connection_timeout_seconds, + ) + self.max_connections = config_dict.get("max_connections", 100) + self.max_keepalive_connections = config_dict.get( + "max_keepalive_connections", + 20, + ) + self.max_concurrent_requests: typing.Optional[int] = config_dict.get( + "max_concurrent_requests", + ) def _handle_nearest_node( self, @@ -295,6 +332,32 @@ def validate_config_dict(config_dict: ConfigDict) -> None: if nearest_node: ConfigurationValidations.validate_nearest_node(nearest_node) + ConfigurationValidations.validate_connection_pool(config_dict) + + @staticmethod + def validate_connection_pool(config_dict: ConfigDict) -> None: + """ + Validate the connection pool and concurrency settings. + + Args: + config_dict (ConfigDict): The configuration dictionary to validate. + + Raises: + ConfigError: If a pool or concurrency setting is out of range. + """ + positive_settings: typing.Dict[str, typing.Optional[float]] = { + "pool_timeout_seconds": config_dict.get("pool_timeout_seconds"), + "max_connections": config_dict.get("max_connections"), + "max_concurrent_requests": config_dict.get("max_concurrent_requests"), + } + for key, config_value in positive_settings.items(): + if config_value is not None and config_value <= 0: + raise ConfigError(f"`{key}` must be greater than 0.") + + max_keepalive_connections = config_dict.get("max_keepalive_connections") + if max_keepalive_connections is not None and max_keepalive_connections < 0: + raise ConfigError("`max_keepalive_connections` must not be negative.") + @staticmethod def validate_required_config_fields(config_dict: ConfigDict) -> None: """ diff --git a/src/typesense/sync/api_call.py b/src/typesense/sync/api_call.py index f65d320..4184599 100644 --- a/src/typesense/sync/api_call.py +++ b/src/typesense/sync/api_call.py @@ -37,6 +37,7 @@ import httpx +from typesense.concurrency_limit import ConcurrencyLimit from typesense.configuration import Configuration, Node from typesense.exceptions import ( HTTPStatus0Error, @@ -174,7 +175,17 @@ def __init__(self, config: Configuration): self.node_manager = NodeManager(config) self.request_handler = RequestHandler(config) self._client = httpx.Client( - timeout=config.connection_timeout_seconds, + timeout=httpx.Timeout( + config.connection_timeout_seconds, + pool=config.pool_timeout_seconds, + ), + limits=httpx.Limits( + max_connections=config.max_connections, + max_keepalive_connections=config.max_keepalive_connections, + ), + ) + self._concurrency_limit = ConcurrencyLimit( + config.max_concurrent_requests, ) def __enter__(self) -> "ApiCall": @@ -519,14 +530,15 @@ def _make_request_and_process_response( **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> typing.Union[TEntityDict, str]: """Make the async API request to `node` and process the response.""" - request_response = self.request_handler.make_request( - method=method, - url=url, - as_json=as_json, - entity_type=entity_type, - client=self._client, - **kwargs, - ) + with self._concurrency_limit: + request_response = self.request_handler.make_request( + method=method, + url=url, + as_json=as_json, + entity_type=entity_type, + client=self._client, + **kwargs, + ) self.node_manager.set_node_health(node, is_healthy=True) return ( typing.cast(TEntityDict, request_response) diff --git a/tests/api_call_test.py b/tests/api_call_test.py index fb4857b..9c77d03 100644 --- a/tests/api_call_test.py +++ b/tests/api_call_test.py @@ -1,8 +1,11 @@ """Unit Tests for the ApiCall class.""" +import asyncio import logging import sys +import threading import time +from concurrent.futures import ThreadPoolExecutor from pytest_mock import MockFixture @@ -812,3 +815,131 @@ def test_success_marks_only_the_answering_node_healthy( assert answering_node.healthy is True assert answering_node.last_access_ts > 0 assert unhealthy_node.healthy is False + + +def test_client_uses_connection_pool_settings( + fake_config: Configuration, + mocker: MockerFixture, +) -> None: + """Test that the httpx client is built from the connection pool settings.""" + client_mock = mocker.patch("typesense.sync.api_call.httpx.Client") + fake_config.connection_timeout_seconds = 3.0 + fake_config.pool_timeout_seconds = 1.5 + fake_config.max_connections = 200 + fake_config.max_keepalive_connections = 50 + + ApiCall(fake_config) + + client_mock.assert_called_once_with( + timeout=httpx.Timeout(3.0, pool=1.5), + limits=httpx.Limits(max_connections=200, max_keepalive_connections=50), + ) + + +def test_async_client_uses_connection_pool_settings( + fake_config: Configuration, + mocker: MockerFixture, +) -> None: + """Test that the httpx async client is built from the connection pool settings.""" + client_mock = mocker.patch("typesense.async_.api_call.httpx.AsyncClient") + fake_config.connection_timeout_seconds = 3.0 + fake_config.pool_timeout_seconds = 1.5 + fake_config.max_connections = 200 + fake_config.max_keepalive_connections = 50 + + AsyncApiCall(fake_config) + + client_mock.assert_called_once_with( + timeout=httpx.Timeout(3.0, pool=1.5), + limits=httpx.Limits(max_connections=200, max_keepalive_connections=50), + ) + + +def _count_requests_in_flight( + concurrent_requests: int, + max_concurrent_requests: typing.Optional[int], + fake_config: Configuration, +) -> int: + """Send requests from several threads and return the peak number in flight.""" + fake_config.max_concurrent_requests = max_concurrent_requests + api_call = ApiCall(fake_config) + lock = threading.Lock() + in_flight = 0 + peak = 0 + + def slow_response(request: httpx.Request) -> httpx.Response: + nonlocal in_flight, peak + with lock: + in_flight += 1 + peak = max(peak, in_flight) + time.sleep(0.05) + with lock: + in_flight -= 1 + return httpx.Response(200, json={"key": "value"}) + + with respx.mock: + respx.get("http://nearest:8108/").mock(side_effect=slow_response) + with ThreadPoolExecutor(max_workers=concurrent_requests) as executor: + for _ in range(concurrent_requests): + executor.submit(api_call.get, "/", entity_type=typing.Dict[str, str]) + + return peak + + +def test_max_concurrent_requests_caps_requests_in_flight( + fake_config: Configuration, +) -> None: + """Test that no more than ``max_concurrent_requests`` requests are in flight.""" + assert _count_requests_in_flight(6, 2, fake_config) == 2 + + +def test_requests_in_flight_are_unlimited_by_default( + fake_config: Configuration, +) -> None: + """Test that requests are not capped when ``max_concurrent_requests`` is unset.""" + assert _count_requests_in_flight(6, None, fake_config) == 6 + + +async def _async_count_requests_in_flight( + concurrent_requests: int, + max_concurrent_requests: typing.Optional[int], + fake_config: Configuration, +) -> int: + """Send concurrent async requests and return the peak number in flight.""" + fake_config.max_concurrent_requests = max_concurrent_requests + api_call = AsyncApiCall(fake_config) + in_flight = 0 + peak = 0 + + async def slow_response(request: httpx.Request) -> httpx.Response: + nonlocal in_flight, peak + in_flight += 1 + peak = max(peak, in_flight) + await asyncio.sleep(0.01) + in_flight -= 1 + return httpx.Response(200, json={"key": "value"}) + + with respx.mock: + respx.get("http://nearest:8108/").mock(side_effect=slow_response) + await asyncio.gather( + *( + api_call.get("/", entity_type=typing.Dict[str, str]) + for _ in range(concurrent_requests) + ), + ) + + return peak + + +async def test_async_max_concurrent_requests_caps_requests_in_flight( + fake_config: Configuration, +) -> None: + """Test that no more than ``max_concurrent_requests`` requests are in flight (async).""" + assert await _async_count_requests_in_flight(6, 2, fake_config) == 2 + + +async def test_async_requests_in_flight_are_unlimited_by_default( + fake_config: Configuration, +) -> None: + """Test that async requests are not capped when ``max_concurrent_requests`` is unset.""" + assert await _async_count_requests_in_flight(6, None, fake_config) == 6 diff --git a/tests/configuration_test.py b/tests/configuration_test.py index 626c477..092c93b 100644 --- a/tests/configuration_test.py +++ b/tests/configuration_test.py @@ -207,3 +207,46 @@ def test_configuration_invalid_nearest_node_url() -> None: match="Node URL does not contain the port.", ): Configuration(config) + + +def test_configuration_connection_pool_defaults() -> None: + """Test the connection pool defaults, with the pool timeout following the connection timeout.""" + configuration = Configuration( + { + "nodes": [DEFAULT_NODE], + "api_key": "xyz", + "connection_timeout_seconds": 7.0, + }, + ) + + expected = { + "pool_timeout_seconds": 7.0, + "max_connections": 100, + "max_keepalive_connections": 20, + "max_concurrent_requests": None, + } + + assert_to_contain_object(configuration, expected) + + +def test_configuration_connection_pool_explicit() -> None: + """Test the connection pool settings with explicit values.""" + configuration = Configuration( + { + "nodes": [DEFAULT_NODE], + "api_key": "xyz", + "pool_timeout_seconds": 1.5, + "max_connections": 200, + "max_keepalive_connections": 50, + "max_concurrent_requests": 150, + }, + ) + + expected = { + "pool_timeout_seconds": 1.5, + "max_connections": 200, + "max_keepalive_connections": 50, + "max_concurrent_requests": 150, + } + + assert_to_contain_object(configuration, expected) diff --git a/tests/configuration_validations_test.py b/tests/configuration_validations_test.py index d408e05..8cf8061 100644 --- a/tests/configuration_validations_test.py +++ b/tests/configuration_validations_test.py @@ -1,7 +1,13 @@ """Tests for the ConfigurationValidations class.""" +import sys import types +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + import pytest from typesense.configuration import ConfigDict, ConfigurationValidations @@ -199,3 +205,34 @@ def test_validate_config_dict_with_wrong_nearest_node() -> None: "api_key": "xyz", }, ) + + +@pytest.mark.parametrize( + ("key", "config_value", "message"), + [ + ("pool_timeout_seconds", 0, "`pool_timeout_seconds` must be greater than 0."), + ("max_connections", 0, "`max_connections` must be greater than 0."), + ( + "max_concurrent_requests", + -1, + "`max_concurrent_requests` must be greater than 0.", + ), + ( + "max_keepalive_connections", + -1, + "`max_keepalive_connections` must not be negative.", + ), + ], +) +def test_validate_config_dict_with_invalid_connection_pool( + key: str, + config_value: float, + message: str, +) -> None: + """Test validate_config_dict with out-of-range connection pool settings.""" + config_dict = {"nodes": [DEFAULT_NODE], "api_key": "xyz", key: config_value} + + with pytest.raises(ConfigError, match=message): + ConfigurationValidations.validate_config_dict( + typing.cast(ConfigDict, config_dict), + ) diff --git a/utils/run-unasync.py b/utils/run-unasync.py index aa4dcbd..7d836d4 100644 --- a/utils/run-unasync.py +++ b/utils/run-unasync.py @@ -27,6 +27,8 @@ def collect_class_replacements(source_dir: Path) -> dict[str, str]: # client (unasync strips ``await``); map the module token so the import and call # are rewritten too. replacements["asyncio"] = "time" + # Defined in the shared ``typesense.concurrency_limit`` module, outside async_. + replacements["AsyncConcurrencyLimit"] = "ConcurrencyLimit" return replacements