diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index eb64f494a..5533efcbe 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -36,6 +36,7 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.tables import TABLES, VIEWS from tests.pyathena.util import succeeded_query_execution, throttle_metadata_api, unreachable_glue _logger = logging.getLogger(__name__) @@ -1571,34 +1572,57 @@ def test_other_metadata_exception_keeps_its_retries(self, cursor, monkeypatch): cursor.get_table_metadata("one_row") assert calls == ["get_table_metadata"] * 2 + # Without the Glue fallback, a throttled page is retried, so Athena serves every page. + @pytest.mark.parametrize("cursor", [{"glue_metadata_fallback": False}], indirect=["cursor"]) + def test_listing_reads_every_page(self, cursor, monkeypatch): + client = cursor.connection.client + list_table_metadata = client.list_table_metadata + tokens = [] + + def record_token(**kwargs): + tokens.append(kwargs.get("NextToken")) + return list_table_metadata(**kwargs) + + monkeypatch.setattr(client, "list_table_metadata", record_token) + # Other tests add and drop their own tables, so only the session's are compared. + session_tables = sorted(t.name for t in (*TABLES, *VIEWS)) + + names = [m.name for m in cursor.list_table_metadata(max_results=4)] + + assert sorted(n for n in names if n in session_tables) == session_tables + # At four per page, the session's tables and views span three pages or more. + assert len(set(tokens)) >= 3 + @pytest.mark.parametrize( "cursor", [{"retry_config": RetryConfig(attempt=1)}], indirect=["cursor"] ) def test_listing_resumes_after_a_failed_glue_request(self, cursor, monkeypatch): # A page throttled mid-listing is read again, not the pages before it. - expected = sorted(m.name for m in cursor.list_table_metadata(max_results=2)) + # The pages are stubbed so that real throttling cannot change the requests. + pages = { + None: {"TableMetadataList": [{"Name": "t0"}, {"Name": "t1"}], "NextToken": "1"}, + "1": {"TableMetadataList": [{"Name": "t2"}, {"Name": "t3"}], "NextToken": "2"}, + "2": {"TableMetadataList": [{"Name": "t4"}]}, + } client = cursor.connection.client - list_table_metadata = client.list_table_metadata requests = [] - def throttle_second_page_once(**kwargs): - requests.append(kwargs.get("NextToken")) + def throttle_second_request_once(**kwargs): + token = kwargs.get("NextToken") + requests.append(token) if len(requests) == 2: raise ClientError( {"Error": {"Code": "ThrottlingException", "Message": "Rate exceeded"}}, "ListTableMetadata", ) - return list_table_metadata(**kwargs) + return pages[token] - monkeypatch.setattr(client, "list_table_metadata", throttle_second_page_once) - self._unreachable_glue(cursor.connection, monkeypatch) + monkeypatch.setattr(client, "list_table_metadata", throttle_second_request_once) + glue = self._unreachable_glue(cursor.connection, monkeypatch) - assert sorted(m.name for m in cursor.list_table_metadata(max_results=2)) == expected - # Every page once, and the throttled second page a second time. - assert len(expected) > 2 - assert requests[0] is None - assert requests[1] == requests[2] - assert len(requests) == len(set(requests)) + 1 + assert [m.name for m in cursor.list_table_metadata()] == ["t0", "t1", "t2", "t3", "t4"] + assert requests == [None, "1", "1", "2"] + assert not glue.reachable @pytest.mark.parametrize( ("cursor", "catalog_name"),