diff --git a/tests/unit/test_client.py b/tests/unit/test_client.py index c5d33dbd..4a56d9af 100644 --- a/tests/unit/test_client.py +++ b/tests/unit/test_client.py @@ -1752,3 +1752,43 @@ def statement_uri(token): assert query.stats["state"] == "FINISHED" # The count row survives draining and the rows stay lazily iterable. assert list(result) == [[3]] + + +def _retryable_get(monkeypatch, body: bytes, status_code: int = 200): + """Serve a canned response from Session.get and return (request, recorder).""" + http_resp = TrinoRequest.http.Response() + http_resp.status_code = status_code + http_resp._content = body + + get_retry = RetryRecorder(result=http_resp) + monkeypatch.setattr(TrinoRequest.http.Session, "get", get_retry) + + req = TrinoRequest( + host="coordinator", + port=8080, + client_session=ClientSession(user="test"), + max_attempts=3, + ) + return req, get_retry + + +def test_empty_body_retry_check_does_not_decode_body(monkeypatch): + # Reading .text would run charset detection over the whole payload merely to test emptiness, + # which is wasteful for a large binary body. The check must look at the raw bytes instead. + req, get_retry = _retryable_get(monkeypatch, body=b"\x89PNG\r\n\x1a\n" * 1024) + + with mock.patch.object( + TrinoRequest.http.Response, "text", new_callable=mock.PropertyMock + ) as response_text: + req.get("URL") + + response_text.assert_not_called() + assert get_retry.retry_count == 1 + + +def test_whitespace_only_200_response_retry(monkeypatch): + req, get_retry = _retryable_get(monkeypatch, body=b" \r\n\t ") + + req.get("URL") + + assert get_retry.retry_count == 3 diff --git a/tests/unit/test_client_spooling.py b/tests/unit/test_client_spooling.py index 21912d37..d5093c76 100644 --- a/tests/unit/test_client_spooling.py +++ b/tests/unit/test_client_spooling.py @@ -13,7 +13,9 @@ import time from unittest import mock +import httpretty import pytest +from httpretty import httprettified from trino.client import _RequestHeartbeat from trino.client import ClientSession @@ -287,7 +289,7 @@ def fake_get(uri, headers=None, **kwargs): recorded["headers"] = headers return mock.Mock(ok=True) - segment._request._get = fake_get + segment._request._get_accept_empty_body = fake_get segment._send_spooling_request(segment.uri) assert recorded["headers"]["X-Auth-Gateway-Token"] == "user-token" @@ -304,7 +306,7 @@ def fake_get(uri, headers=None, **kwargs): recorded["headers"] = headers return mock.Mock(ok=True) - segment._request._get = fake_get + segment._request._get_accept_empty_body = fake_get external_uri = "https://s3.amazonaws.com/bucket/seg1?X-Amz-Signature=abc" segment._send_spooling_request(external_uri) @@ -323,7 +325,44 @@ def fake_get(uri, headers=None, **kwargs): recorded["headers"] = headers return mock.Mock(ok=True) - segment._request._get = fake_get + segment._request._get_accept_empty_body = fake_get segment._send_spooling_request(segment.uri) assert recorded["headers"]["X-Trino-Spooling-Token"] == "token-abc" + + +def _spooled_segment_for_ack(max_attempts): + request = TrinoRequest( + host="coordinator", + port=8080, + client_session=ClientSession(user="test"), + http_scheme="http", + max_attempts=max_attempts, + ) + segment_to = { + "type": "spooled", + "uri": "http://coordinator/v1/spooled/download/seg1", + "ackUri": "http://coordinator/v1/spooled/ack/seg1", + "metadata": {"segmentSize": "1", "uncompressedSize": "1"}, + } + return SpooledSegment(segment_to, request) + + +@httprettified +def test_acknowledge_request_retries_on_error_status_code(): + segment = _spooled_segment_for_ack(max_attempts=3) + httpretty.register_uri(method=httpretty.GET, uri=segment.ack_uri, body="", status=503) + + segment._send_acknowledgement() + + assert len(httpretty.latest_requests()) == 3 + + +@httprettified +def test_acknowledge_request_does_not_retry_on_empty_200_response(): + segment = _spooled_segment_for_ack(max_attempts=3) + httpretty.register_uri(method=httpretty.GET, uri=segment.ack_uri, body="", status=200) + + segment._send_acknowledgement() + + assert len(httpretty.latest_requests()) == 1 diff --git a/trino/client.py b/trino/client.py index 9d4956bf..753fe27c 100644 --- a/trino/client.py +++ b/trino/client.py @@ -636,29 +636,36 @@ def max_attempts(self) -> int: def max_attempts(self, value: int) -> None: self._max_attempts = value if value == 1: # No retry - self._get = self._http_session.get + self._get = self._get_accept_empty_body = self._http_session.get self._post = self._http_session.post self._delete = self._http_session.delete self._head = self._http_session.head return + def has_error_status(response: Response) -> bool: + return getattr(response, "status_code", None) in (429, 502, 503, 504) + + def has_ok_status_but_no_body(response: Response) -> bool: + return getattr(response, "status_code", None) == 200 and not getattr(response, "content", b"").strip() + with_retry = _retry_with( self._handle_retry, handled_exceptions=self._exceptions, - conditions=( - # need retry when there is no exception but the status code is 429, 502, 503, or 504 - lambda response: getattr(response, "status_code", None) - in (429, 502, 503, 504), - # need retry when the server returns 200 with an empty body (transient under load) - lambda response: getattr(response, "status_code", None) == 200 - and not getattr(response, "text", "").strip(), - ), + # retry when there is no exception but error status_code, and when status code is 200 but there's no body + conditions=(has_error_status, has_ok_status_but_no_body), max_attempts=self._max_attempts, ) + self._get = with_retry(self._http_session.get) self._post = with_retry(self._http_session.post) self._delete = with_retry(self._http_session.delete) self._head = with_retry(self._http_session.head) + self._get_accept_empty_body = _retry_with( + self._handle_retry, + handled_exceptions=self._exceptions, + conditions=(has_error_status,), + max_attempts=self._max_attempts, + )(self._http_session.get) def get_url(self, path: str) -> str: return "{protocol}://{host}:{port}{path}".format( @@ -1310,15 +1317,16 @@ def headers(self) -> Dict[str, List[str]]: return self._segment.get("headers", {}) def acknowledge(self) -> None: - def acknowledge_request(): - try: - http_response = self._send_spooling_request(self.ack_uri, timeout=2) - if not http_response.ok: - self._request.raise_response_error(http_response) - except Exception as e: - logger.error(f"Failed to acknowledge spooling request for segment {self}: {e}") # Start the acknowledgment in the executor thread - executor.submit(acknowledge_request) + executor.submit(self._send_acknowledgement) + + def _send_acknowledgement(self): + try: + http_response = self._send_spooling_request(self.ack_uri, timeout=2) + if not http_response.ok: + self._request.raise_response_error(http_response) + except Exception as e: + logger.error(f"Failed to acknowledge spooling request for segment {self}: {e}") def _send_spooling_request(self, uri: str, **kwargs) -> requests.Response: headers: Dict[str, str] = {} @@ -1332,7 +1340,7 @@ def _send_spooling_request(self, uri: str, **kwargs) -> requests.Response: if len(values) > 1: raise ValueError(f"Header '{key}' contains multiple values: {values}") headers[key] = values[0] - return self._request._get(uri, headers=headers, **kwargs) + return self._request._get_accept_empty_body(uri, headers=headers, **kwargs) def __repr__(self): return (