Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 41 additions & 14 deletions src/typesense/async_/api_call.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@

import httpx

from typesense.concurrency_limit import AsyncConcurrencyLimit
from typesense.configuration import Configuration, Node
from typesense.exceptions import (
HTTPStatus0Error,
Expand Down Expand Up @@ -135,6 +136,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:
"""
Expand All @@ -160,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":
Expand Down Expand Up @@ -473,11 +498,14 @@ async def _execute_request(
try:
return await self._make_request_and_process_response(
method,
node,
url,
entity_type,
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:
Expand All @@ -495,24 +523,23 @@ 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."""
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(
self.node_manager.get_node(),
is_healthy=True,
)
"""Make the async API request to `node` and process the response."""
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)
if as_json
Expand Down
89 changes: 89 additions & 0 deletions src/typesense/concurrency_limit.py
Original file line number Diff line number Diff line change
@@ -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()
63 changes: 63 additions & 0 deletions src/typesense/configuration.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]]
Expand All @@ -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:
Expand Down Expand Up @@ -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__(
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
"""
Expand Down
55 changes: 41 additions & 14 deletions src/typesense/sync/api_call.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@

import httpx

from typesense.concurrency_limit import ConcurrencyLimit
from typesense.configuration import Configuration, Node
from typesense.exceptions import (
HTTPStatus0Error,
Expand Down Expand Up @@ -135,6 +136,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:
"""
Expand All @@ -160,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":
Expand Down Expand Up @@ -473,11 +498,14 @@ def _execute_request(
try:
return self._make_request_and_process_response(
method,
node,
url,
entity_type,
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:
Expand All @@ -495,24 +523,23 @@ 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."""
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(
self.node_manager.get_node(),
is_healthy=True,
)
"""Make the async API request to `node` and process the response."""
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)
if as_json
Expand Down
Loading
Loading