Skip to content
Draft
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
21 changes: 21 additions & 0 deletions src/typesense/async_/api_call.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
by other components of the library.
"""

import asyncio
import sys
from types import MappingProxyType, TracebackType

Expand Down Expand Up @@ -134,6 +135,22 @@ class SessionFunctionKwargs(typing.Generic[TParams, TBody], typing.TypedDict):
ServiceUnavailable,
)

# Raised by httpx inside the client, so they say nothing about the node's
# health. They subclass entries of _SERVER_ERRORS and must be caught first.
_CLIENT_SIDE_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 Down Expand Up @@ -477,8 +494,12 @@ async def _execute_request(
as_json,
**request_kwargs,
)
except _CLIENT_SIDE_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:
await asyncio.sleep(self.config.retry_interval_seconds)
return await self._execute_request(
method,
endpoint,
Expand Down
14 changes: 11 additions & 3 deletions src/typesense/configuration.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,9 @@ class ConfigDict(typing.TypedDict):

num_retries (int): The number of retries to attempt before failing.

interval_seconds (int): The interval in seconds between retries.
retry_interval_seconds (float): The interval in seconds between retries.

interval_seconds (int): Deprecated alias of ``retry_interval_seconds``.

healthcheck_interval_seconds (int): The interval in seconds between
health checks.
Expand All @@ -86,7 +88,8 @@ class ConfigDict(typing.TypedDict):
nearest_node: typing.NotRequired[typing.Union[str, NodeConfigDict]]
api_key: str
num_retries: typing.NotRequired[int]
interval_seconds: typing.NotRequired[int]
retry_interval_seconds: typing.NotRequired[float]
interval_seconds: typing.NotRequired[int] # deprecated alias
healthcheck_interval_seconds: typing.NotRequired[int]
verify: typing.NotRequired[bool]
timeout_seconds: typing.NotRequired[int] # deprecated
Expand Down Expand Up @@ -214,7 +217,12 @@ def __init__(
3.0,
)
self.num_retries = config_dict.get("num_retries", 3)
self.retry_interval_seconds = config_dict.get("retry_interval_seconds", 1.0)
# ``interval_seconds`` is the historically documented key; ``retry_interval_seconds``
# is what this attribute is named. Honor both so the documented spelling works too.
self.retry_interval_seconds = config_dict.get(
"retry_interval_seconds",
config_dict.get("interval_seconds", 1.0),
)
self.healthcheck_interval_seconds = config_dict.get(
"healthcheck_interval_seconds",
60,
Expand Down
21 changes: 21 additions & 0 deletions src/typesense/sync/api_call.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
by other components of the library.
"""

import time
import sys
from types import MappingProxyType, TracebackType

Expand Down Expand Up @@ -134,6 +135,22 @@ class SessionFunctionKwargs(typing.Generic[TParams, TBody], typing.TypedDict):
ServiceUnavailable,
)

# Raised by httpx inside the client, so they say nothing about the node's
# health. They subclass entries of _SERVER_ERRORS and must be caught first.
_CLIENT_SIDE_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 Down Expand Up @@ -477,8 +494,12 @@ def _execute_request(
as_json,
**request_kwargs,
)
except _CLIENT_SIDE_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:
time.sleep(self.config.retry_interval_seconds)
return self._execute_request(
method,
endpoint,
Expand Down
108 changes: 108 additions & 0 deletions tests/api_call_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
from tests.utils.object_assertions import assert_match_object, assert_object_lists_match
from typesense import exceptions
from typesense.sync.api_call import ApiCall, RequestHandler
from typesense.async_.api_call import AsyncApiCall
from typesense.configuration import Configuration, Node
from typesense.logger import logger

Expand Down Expand Up @@ -615,3 +616,110 @@ def test_max_retries_no_last_exception(fake_api_call: ApiCall) -> None:
num_retries=10,
last_exception=None,
)


def test_sleeps_retry_interval_between_retries(
fake_api_call: ApiCall,
mocker: MockerFixture,
) -> None:
"""Test that it waits ``retry_interval_seconds`` between failed attempts."""
sleep_mock = mocker.patch("typesense.sync.api_call.time.sleep")

with respx.mock:
for host in ("nearest", "node0", "node1", "node2"):
respx.get(f"http://{host}:8108/").mock(
return_value=httpx.Response(503, json={"message": "unavailable"}),
)

with pytest.raises(exceptions.ServiceUnavailable):
fake_api_call.get("/", entity_type=typing.Dict[str, str])

# ``num_retries`` gaps for ``num_retries + 1`` attempts, and each gap must be
# ``retry_interval_seconds`` long (regression: the delay was dropped entirely).
assert sleep_mock.call_count == fake_api_call.config.num_retries
for sleep_call in sleep_mock.call_args_list:
assert sleep_call == mocker.call(fake_api_call.config.retry_interval_seconds)


async def test_async_sleeps_retry_interval_between_retries(
fake_async_api_call: AsyncApiCall,
mocker: MockerFixture,
) -> None:
"""Test that the async client waits ``retry_interval_seconds`` between attempts."""
sleep_mock = mocker.patch(
"typesense.async_.api_call.asyncio.sleep",
new_callable=mocker.AsyncMock,
)

with respx.mock:
for host in ("nearest", "node0", "node1", "node2"):
respx.get(f"http://{host}:8108/").mock(
return_value=httpx.Response(503, json={"message": "unavailable"}),
)

with pytest.raises(exceptions.ServiceUnavailable):
await fake_async_api_call.get("/", entity_type=typing.Dict[str, str])

assert sleep_mock.call_count == fake_async_api_call.config.num_retries
for sleep_call in sleep_mock.call_args_list:
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
4 changes: 4 additions & 0 deletions utils/run-unasync.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,10 @@ def collect_class_replacements(source_dir: Path) -> dict[str, str]:
async_name = match.group(1)
replacements[async_name] = async_name[len("Async") :]
replacements["aclose"] = "close"
# ``await asyncio.sleep`` in the async client becomes ``time.sleep`` in the sync
# client (unasync strips ``await``); map the module token so the import and call
# are rewritten too.
replacements["asyncio"] = "time"
return replacements


Expand Down
Loading