Skip to content
Closed
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
40 changes: 40 additions & 0 deletions tests/unit/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
45 changes: 42 additions & 3 deletions tests/unit/test_client_spooling.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"
Expand All @@ -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)

Expand All @@ -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
44 changes: 26 additions & 18 deletions trino/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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] = {}
Expand All @@ -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 (
Expand Down
Loading