diff --git a/src/typesense/request_handler.py b/src/typesense/request_handler.py index 38e6c24..9b49fea 100644 --- a/src/typesense/request_handler.py +++ b/src/typesense/request_handler.py @@ -309,10 +309,15 @@ def _get_exception(http_code: int) -> typing.Type[TypesenseClientError]: """ Map an HTTP status code to the appropriate exception type. + Any 5xx code without a dedicated exception (e.g. 502, 504 from a proxy + in front of a restarting node) maps to ServerError, so the request is + retried on another node. + Args: http_code (int): The HTTP status code. Returns: Type[TypesenseClientError]: The exception type corresponding to the status code. """ - return _ERROR_CODE_MAP.get(str(http_code), TypesenseClientError) + default = ServerError if 500 <= http_code <= 599 else TypesenseClientError + return _ERROR_CODE_MAP.get(str(http_code), default) diff --git a/tests/api_call_test.py b/tests/api_call_test.py index ddff4ee..bdc1386 100644 --- a/tests/api_call_test.py +++ b/tests/api_call_test.py @@ -88,6 +88,11 @@ def test_get_exception() -> None: assert RequestHandler._get_exception(422) == exceptions.ObjectUnprocessable assert RequestHandler._get_exception(500) == exceptions.ServerError assert RequestHandler._get_exception(503) == exceptions.ServiceUnavailable + assert RequestHandler._get_exception(501) == exceptions.ServerError + assert RequestHandler._get_exception(502) == exceptions.ServerError + assert RequestHandler._get_exception(504) == exceptions.ServerError + assert RequestHandler._get_exception(599) == exceptions.ServerError + assert RequestHandler._get_exception(418) == exceptions.TypesenseClientError assert RequestHandler._get_exception(999) == exceptions.TypesenseClientError @@ -460,6 +465,33 @@ def test_selects_next_available_node_on_timeout( assert len(respx.calls) == 3 +@pytest.mark.parametrize("status_code", [500, 502, 503, 504]) +def test_selects_next_available_node_on_5xx( + fake_api_call: ApiCall, + status_code: int, +) -> None: + """Test that a 5xx response marks the node unhealthy and retries on the next one.""" + with respx.mock: + respx.get("http://nearest:8108/test").mock( + return_value=httpx.Response(status_code, text="Bad Gateway") + ) + respx.get("http://node0:8108/test").mock( + return_value=httpx.Response(200, json={"key": "value"}) + ) + + response = fake_api_call.get( + "/test", + as_json=True, + entity_type=typing.Dict[str, str], + ) + + assert response == {"key": "value"} + assert respx.calls[0].request.url == "http://nearest:8108/test" + assert respx.calls[1].request.url == "http://node0:8108/test" + assert len(respx.calls) == 2 + assert fake_api_call.config.nearest_node.healthy is False + + def test_get_node_no_healthy_nodes( fake_api_call: ApiCall, mocker: MockFixture,