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
25 changes: 20 additions & 5 deletions src/typesense/async_/api_call.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
"""
Expand Down Expand Up @@ -473,11 +487,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,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,
Expand All @@ -509,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
Expand Down
25 changes: 20 additions & 5 deletions src/typesense/sync/api_call.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
"""
Expand Down Expand Up @@ -473,11 +487,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,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,
Expand All @@ -509,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
Expand Down
147 changes: 147 additions & 0 deletions tests/api_call_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -665,3 +684,131 @@ 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


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
Loading