diff --git a/docs/testing.md b/docs/testing.md index bf58fb77..41addc5c 100644 --- a/docs/testing.md +++ b/docs/testing.md @@ -107,6 +107,14 @@ The test session hooks upload fixture data to S3 and create and remove Athena/Gl Selecting a single test under `tests/pyathena/` still invokes the session setup, even when the test itself uses only mocks or pure Python logic. Do not assume that `-k` or a single test path makes this suite offline. +The tables and views from `tests/pyathena/tables.py` and the data files they read are in a fixture schema, `ENV.fixture_schema`, which tests only read. +With pytest-xdist, the controller creates it once before the workers start and drops it after they finish. +A run without workers creates and drops the fixture schema itself. +The session hooks run only when a path given to pytest is `tests/pyathena/` or inside it, as with `just test pyathena`. +Each test process also creates its own schema, `ENV.schema`, where tests create their tables and views. +The fixture schema is the default schema of the cursor and engine fixtures, so a test that creates an object qualifies its name with `ENV.schema`. +At the end of the run, the process that created the fixture schema lists its tables and fails the run if they differ from the ones it created. + Run the suites relevant to the change: ```bash diff --git a/scripts/sweep_databases.py b/scripts/sweep_databases.py index 1fbf911f..682dc6d6 100644 --- a/scripts/sweep_databases.py +++ b/scripts/sweep_databases.py @@ -71,7 +71,7 @@ def _eligible(database: dict[str, Any], cutoff: datetime) -> bool: def sweep_databases(client: Any, catalog_id: str, *, dry_run: bool = True) -> dict[str, int]: """Preview or delete test databases older than seven days. - Fixtures generate fresh database names for each session or worker. + Fixtures generate fresh database names for each pytest run and each test process. Databases younger than seven days are retained, including concurrent CI runs. Only Glue metadata is deleted; S3 objects and child catalogs are untouched. """ diff --git a/tests/__init__.py b/tests/__init__.py index 124fe300..3c22a684 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -12,6 +12,15 @@ ) +def _random_schema(): + """Return a new test schema name, the shape scripts/sweep_databases.py sweeps. + + Returns: + ``pyathena_test_`` followed by 10 random lowercase letters and digits. + """ + return "pyathena_test_" + "".join(random.choices(string.ascii_lowercase + string.digits, k=10)) + + class Env: def __init__(self): self.region_name = os.getenv("AWS_DEFAULT_REGION") @@ -31,19 +40,29 @@ def __init__(self): ) self.default_work_group = os.getenv("AWS_ATHENA_DEFAULT_WORKGROUP", "primary") self.managed_work_group = os.getenv("AWS_ATHENA_MANAGED_WORKGROUP") - self.schema = "pyathena_test_" + "".join( - random.choices(string.ascii_lowercase + string.digits, k=10) - ) + # Each test process creates the objects of its tests in its own schema. + self.schema = _random_schema() + # The read-only tables and views from tests/pyathena/tables.py, and the + # data files they read, are in their own schema. With pytest-xdist, the + # controller creates it once and passes its name to the workers, which + # replace this value in pytest_configure (tests/pyathena/conftest.py). + self.fixture_schema = _random_schema() # Optional Amazon S3 Tables configuration. `s3tables_catalog` is the # registered table-bucket catalog, e.g. "s3tablescatalog/". - # The test session creates its own namespace in it, named like the schema, + # Each test process creates its own namespace in it, named like the schema, # because a table dropped during a listing of a shared namespace fails # the listing. The S3 Tables tests skip when the catalog is unset. self.s3tables_catalog = os.getenv("AWS_ATHENA_S3_TABLES_CATALOG") self.s3tables_namespace = self.schema if self.s3tables_catalog else None - self.s3_filesystem_test_file_key = ( - f"{self.s3_staging_key}{self.schema}/filesystem/test_read/test.dat" - ) + + @property + def s3_filesystem_test_file_key(self): + """The S3 key of the read-only file the filesystem tests read. + + Returns: + The key, under the fixture schema's prefix in the staging directory. + """ + return f"{self.s3_staging_key}{self.fixture_schema}/filesystem/test_read/test.dat" ENV = Env() diff --git a/tests/pyathena/aio/arrow/test_cursor.py b/tests/pyathena/aio/arrow/test_cursor.py index 9346d3bd..cb427f5f 100644 --- a/tests/pyathena/aio/arrow/test_cursor.py +++ b/tests/pyathena/aio/arrow/test_cursor.py @@ -80,7 +80,7 @@ async def test_no_result_set_raises(self, aio_arrow_cursor): async def test_context_manager(self): from pyathena.aio.arrow.cursor import AioArrowCursor - conn = await _aio_connect(schema_name=ENV.schema, cursor_class=AioArrowCursor) + conn = await _aio_connect(schema_name=ENV.fixture_schema, cursor_class=AioArrowCursor) try: async with conn.cursor() as cursor: await cursor.execute("SELECT * FROM one_row") diff --git a/tests/pyathena/aio/conftest.py b/tests/pyathena/aio/conftest.py index 61ec80bd..6d7e1a24 100644 --- a/tests/pyathena/aio/conftest.py +++ b/tests/pyathena/aio/conftest.py @@ -20,11 +20,21 @@ async def _aio_connect(schema_name="default", **kwargs): @pytest.fixture async def aio_cursor(request): + """Yield an ``AioCursor`` whose default schema is the fixture schema. + + Args: + request: The fixture request; its optional ``param`` holds connection options. + + Yields: + The cursor. + """ from pyathena.aio.cursor import AioCursor if not hasattr(request, "param"): request.param = {} - conn = await _aio_connect(schema_name=ENV.schema, cursor_class=AioCursor, **request.param) + conn = await _aio_connect( + schema_name=ENV.fixture_schema, cursor_class=AioCursor, **request.param + ) try: async with conn.cursor() as cursor: yield cursor @@ -34,11 +44,21 @@ async def aio_cursor(request): @pytest.fixture async def aio_dict_cursor(request): + """Yield an ``AioDictCursor`` whose default schema is the fixture schema. + + Args: + request: The fixture request; its optional ``param`` holds connection options. + + Yields: + The cursor. + """ from pyathena.aio.cursor import AioDictCursor if not hasattr(request, "param"): request.param = {} - conn = await _aio_connect(schema_name=ENV.schema, cursor_class=AioDictCursor, **request.param) + conn = await _aio_connect( + schema_name=ENV.fixture_schema, cursor_class=AioDictCursor, **request.param + ) try: async with conn.cursor() as cursor: yield cursor @@ -48,11 +68,21 @@ async def aio_dict_cursor(request): @pytest.fixture async def aio_pandas_cursor(request): + """Yield an ``AioPandasCursor`` whose default schema is the fixture schema. + + Args: + request: The fixture request; its optional ``param`` holds connection options. + + Yields: + The cursor. + """ from pyathena.aio.pandas.cursor import AioPandasCursor if not hasattr(request, "param"): request.param = {} - conn = await _aio_connect(schema_name=ENV.schema, cursor_class=AioPandasCursor, **request.param) + conn = await _aio_connect( + schema_name=ENV.fixture_schema, cursor_class=AioPandasCursor, **request.param + ) try: async with conn.cursor() as cursor: yield cursor @@ -62,11 +92,21 @@ async def aio_pandas_cursor(request): @pytest.fixture async def aio_arrow_cursor(request): + """Yield an ``AioArrowCursor`` whose default schema is the fixture schema. + + Args: + request: The fixture request; its optional ``param`` holds connection options. + + Yields: + The cursor. + """ from pyathena.aio.arrow.cursor import AioArrowCursor if not hasattr(request, "param"): request.param = {} - conn = await _aio_connect(schema_name=ENV.schema, cursor_class=AioArrowCursor, **request.param) + conn = await _aio_connect( + schema_name=ENV.fixture_schema, cursor_class=AioArrowCursor, **request.param + ) try: async with conn.cursor() as cursor: yield cursor @@ -76,11 +116,21 @@ async def aio_arrow_cursor(request): @pytest.fixture async def aio_polars_cursor(request): + """Yield an ``AioPolarsCursor`` whose default schema is the fixture schema. + + Args: + request: The fixture request; its optional ``param`` holds connection options. + + Yields: + The cursor. + """ from pyathena.aio.polars.cursor import AioPolarsCursor if not hasattr(request, "param"): request.param = {} - conn = await _aio_connect(schema_name=ENV.schema, cursor_class=AioPolarsCursor, **request.param) + conn = await _aio_connect( + schema_name=ENV.fixture_schema, cursor_class=AioPolarsCursor, **request.param + ) try: async with conn.cursor() as cursor: yield cursor @@ -90,11 +140,21 @@ async def aio_polars_cursor(request): @pytest.fixture async def aio_s3fs_cursor(request): + """Yield an ``AioS3FSCursor`` whose default schema is the fixture schema. + + Args: + request: The fixture request; its optional ``param`` holds connection options. + + Yields: + The cursor. + """ from pyathena.aio.s3fs.cursor import AioS3FSCursor if not hasattr(request, "param"): request.param = {} - conn = await _aio_connect(schema_name=ENV.schema, cursor_class=AioS3FSCursor, **request.param) + conn = await _aio_connect( + schema_name=ENV.fixture_schema, cursor_class=AioS3FSCursor, **request.param + ) try: async with conn.cursor() as cursor: yield cursor @@ -104,6 +164,14 @@ async def aio_s3fs_cursor(request): @pytest.fixture async def aio_spark_cursor(request): + """Yield an ``AioSparkCursor`` whose default schema is the fixture schema. + + Args: + request: The fixture request; its optional ``param`` holds connection options. + + Yields: + The cursor. + """ import asyncio from pyathena.aio.spark.cursor import AioSparkCursor @@ -111,7 +179,7 @@ async def aio_spark_cursor(request): if not hasattr(request, "param"): request.param = {} conn = await _aio_connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=AioSparkCursor, work_group=ENV.spark_work_group, **request.param, diff --git a/tests/pyathena/aio/pandas/test_cursor.py b/tests/pyathena/aio/pandas/test_cursor.py index 4d04c9ae..91dfd749 100644 --- a/tests/pyathena/aio/pandas/test_cursor.py +++ b/tests/pyathena/aio/pandas/test_cursor.py @@ -71,7 +71,7 @@ async def test_no_result_set_raises(self, aio_pandas_cursor): async def test_context_manager(self): from pyathena.aio.pandas.cursor import AioPandasCursor - conn = await _aio_connect(schema_name=ENV.schema, cursor_class=AioPandasCursor) + conn = await _aio_connect(schema_name=ENV.fixture_schema, cursor_class=AioPandasCursor) try: async with conn.cursor() as cursor: await cursor.execute("SELECT * FROM one_row") diff --git a/tests/pyathena/aio/polars/test_cursor.py b/tests/pyathena/aio/polars/test_cursor.py index 41b556a3..2bf0e8d6 100644 --- a/tests/pyathena/aio/polars/test_cursor.py +++ b/tests/pyathena/aio/polars/test_cursor.py @@ -63,7 +63,7 @@ async def test_no_result_set_raises(self, aio_polars_cursor): async def test_context_manager(self): from pyathena.aio.polars.cursor import AioPolarsCursor - conn = await _aio_connect(schema_name=ENV.schema, cursor_class=AioPolarsCursor) + conn = await _aio_connect(schema_name=ENV.fixture_schema, cursor_class=AioPolarsCursor) try: async with conn.cursor() as cursor: await cursor.execute("SELECT * FROM one_row") diff --git a/tests/pyathena/aio/s3fs/test_cursor.py b/tests/pyathena/aio/s3fs/test_cursor.py index 4005da89..add99500 100644 --- a/tests/pyathena/aio/s3fs/test_cursor.py +++ b/tests/pyathena/aio/s3fs/test_cursor.py @@ -71,7 +71,7 @@ async def test_async_iterator(self, aio_s3fs_cursor): assert rows == [(1,)] async def test_context_manager(self): - conn = await _aio_connect(schema_name=ENV.schema) + conn = await _aio_connect(schema_name=ENV.fixture_schema) try: async with conn.cursor(AioS3FSCursor) as cursor: await cursor.execute("SELECT * FROM one_row") diff --git a/tests/pyathena/aio/spark/test_cursor.py b/tests/pyathena/aio/spark/test_cursor.py index 92053ef5..6c211379 100644 --- a/tests/pyathena/aio/spark/test_cursor.py +++ b/tests/pyathena/aio/spark/test_cursor.py @@ -112,7 +112,7 @@ async def test_spark_dataframe(self, aio_spark_cursor): df = spark.read.format("csv") \\ .option("header", "true") \\ .option("inferSchema", "true") \\ - .load("{ENV.s3_staging_dir}{ENV.schema}/spark_group_by/spark_group_by.csv") + .load("{ENV.s3_staging_dir}{ENV.fixture_schema}/spark_group_by/spark_group_by.csv") """ ), description="test description", @@ -172,7 +172,7 @@ async def test_spark_sql(self, aio_spark_cursor): await aio_spark_cursor.execute( textwrap.dedent( f""" - spark.sql("SELECT * FROM {ENV.schema}.one_row").show() + spark.sql("SELECT * FROM {ENV.fixture_schema}.one_row").show() """ ) ) @@ -484,7 +484,7 @@ async def test_executemany(self, aio_spark_cursor): async def test_context_manager(self): conn = await _aio_connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=AioSparkCursor, work_group=ENV.spark_work_group, ) diff --git a/tests/pyathena/aio/sqlalchemy/test_base.py b/tests/pyathena/aio/sqlalchemy/test_base.py index d1d2e614..016a17fe 100644 --- a/tests/pyathena/aio/sqlalchemy/test_base.py +++ b/tests/pyathena/aio/sqlalchemy/test_base.py @@ -105,7 +105,9 @@ async def test_unicode(self, async_engine): async def test_reflect_table(self, async_engine): _, conn = async_engine one_row = await conn.run_sync( - lambda sync_conn: Table("one_row", MetaData(schema=ENV.schema), autoload_with=sync_conn) + lambda sync_conn: Table( + "one_row", MetaData(schema=ENV.fixture_schema), autoload_with=sync_conn + ) ) assert len(one_row.c) == 1 assert one_row.c.number_of_rows is not None @@ -119,7 +121,7 @@ def _inspect(sync_conn): return insp.get_schema_names() schemas = await conn.run_sync(_inspect) - assert ENV.schema in schemas + assert ENV.fixture_schema in schemas assert "default" in schemas async def test_get_table_names(self, async_engine): @@ -127,7 +129,7 @@ async def test_get_table_names(self, async_engine): def _inspect(sync_conn): insp = sqlalchemy.inspect(sync_conn) - return insp.get_table_names(schema=ENV.schema) + return insp.get_table_names(schema=ENV.fixture_schema) table_names = await conn.run_sync(_inspect) assert "many_rows" in table_names @@ -138,9 +140,9 @@ async def test_throttled_reflection_reads_glue(self, async_engine, monkeypatch): def reflect(sync_conn): insp = sqlalchemy.inspect(sync_conn) return ( - insp.get_table_comment("one_row", schema=ENV.schema), - insp.get_table_options("one_row", schema=ENV.schema), - insp.get_table_names(schema=ENV.schema), + insp.get_table_comment("one_row", schema=ENV.fixture_schema), + insp.get_table_options("one_row", schema=ENV.fixture_schema), + insp.get_table_names(schema=ENV.fixture_schema), ) expected = await conn.run_sync(reflect) @@ -160,8 +162,8 @@ async def test_has_table(self, async_engine): def _inspect(sync_conn): insp = sqlalchemy.inspect(sync_conn) return ( - insp.has_table("one_row", schema=ENV.schema), - insp.has_table("this_table_does_not_exist", schema=ENV.schema), + insp.has_table("one_row", schema=ENV.fixture_schema), + insp.has_table("this_table_does_not_exist", schema=ENV.fixture_schema), ) exists, not_exists = await conn.run_sync(_inspect) @@ -173,7 +175,7 @@ async def test_get_columns(self, async_engine): def _inspect(sync_conn): insp = sqlalchemy.inspect(sync_conn) - return insp.get_columns(table_name="one_row", schema=ENV.schema) + return insp.get_columns(table_name="one_row", schema=ENV.fixture_schema) columns = await conn.run_sync(_inspect) actual = columns[0] diff --git a/tests/pyathena/aio/test_cursor.py b/tests/pyathena/aio/test_cursor.py index a1259eff..7f4a1f75 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -44,7 +44,7 @@ async def test_fetchone(self, aio_cursor): assert await aio_cursor.fetchone() == (1,) assert aio_cursor.rownumber == 1 assert await aio_cursor.fetchone() is None - assert aio_cursor.database == ENV.schema + assert aio_cursor.database == ENV.fixture_schema assert aio_cursor.catalog assert aio_cursor.query_id assert aio_cursor.query @@ -440,7 +440,7 @@ async def test_executemany_fetch(self, aio_cursor): await aio_cursor.fetchone() async def test_context_manager(self): - conn = await _aio_connect(schema_name=ENV.schema) + conn = await _aio_connect(schema_name=ENV.fixture_schema) try: async with conn.cursor() as cursor: await cursor.execute("SELECT * FROM one_row") @@ -525,7 +525,8 @@ async def read(): return ( view(await aio_cursor.get_table_metadata("one_row")), sorted(view(m) for m in await aio_cursor.list_table_metadata()), - ENV.schema in [d.name for d in await aio_cursor.list_databases("AwsDataCatalog")], + ENV.fixture_schema + in [d.name for d in await aio_cursor.list_databases("AwsDataCatalog")], ) expected = await read() diff --git a/tests/pyathena/arrow/test_async_cursor.py b/tests/pyathena/arrow/test_async_cursor.py index 734961e3..3be58aea 100644 --- a/tests/pyathena/arrow/test_async_cursor.py +++ b/tests/pyathena/arrow/test_async_cursor.py @@ -136,7 +136,7 @@ def test_query_execution(self, async_arrow_cursor): future = async_arrow_cursor.query_execution(query_id) query_execution = future.result() - assert query_execution.database == ENV.schema + assert query_execution.database == ENV.fixture_schema assert query_execution.catalog assert query_execution.query_id if async_arrow_cursor._unload: diff --git a/tests/pyathena/conftest.py b/tests/pyathena/conftest.py index b63ae25b..68ab5862 100644 --- a/tests/pyathena/conftest.py +++ b/tests/pyathena/conftest.py @@ -1,46 +1,115 @@ import contextlib import functools +import sys import uuid import boto3 import pytest import sqlalchemy +from botocore.exceptions import ClientError from sqlalchemy.ext.asyncio import create_async_engine as _create_async_engine from tests import ASYNC_SQLALCHEMY_CONNECTION_STRING, ENV, SQLALCHEMY_CONNECTION_STRING from tests.pyathena.tables import TABLES, VIEWS, spark_group_by_csv from tests.pyathena.util import read_query +# The key of ENV.fixture_schema in a pytest-xdist worker's workerinput. +_FIXTURE_SCHEMA_KEY = "pyathena_fixture_schema" + +class _XDistHooks: + """pytest-xdist hooks, registered only when the plugin is present.""" + + def pytest_configure_node(self, node): + """Pass the fixture schema's name to a pytest-xdist worker. + + Args: + node: The controller's handle of the worker. + """ + node.workerinput[_FIXTURE_SCHEMA_KEY] = ENV.fixture_schema + + +def pytest_configure(config): + """Share one fixture schema between the pytest-xdist controller and its workers. + + A worker takes the name the controller passed; pytest-xdist sets + ``workerinput`` before it configures the worker. + + Args: + config: The pytest config. + """ + workerinput = getattr(config, "workerinput", None) + if workerinput is not None: + ENV.fixture_schema = workerinput[_FIXTURE_SCHEMA_KEY] + elif config.pluginmanager.hasplugin("xdist"): + config.pluginmanager.register(_XDistHooks()) + + +# The removals of what pytest_sessionstart created, in creation order. +_cleanups = [] + + +@pytest.hookimpl(wrapper=True) def pytest_sessionstart(session): - # The pytest-xdist controller runs no tests, so it sets up nothing. - if not _is_test_process(session.config): - return - _create_s3tables_namespace() - # pytest skips pytest_sessionfinish after a failed pytest_sessionstart, so - # a failure after the namespace is created deletes it here. + """Create the fixture schema and this process's own schema, as its role requires. + + The pytest-xdist controller creates the fixture schema before its own + session start starts the workers, and each worker creates its own schema. + A run without workers creates both. Each removal is recorded before its + step, and ``pytest_sessionfinish`` runs them. pytest skips + ``pytest_sessionfinish`` after a failed session start, so a failure here or + in a later session-start hook runs them at once, before the error reaches + pytest-xdist, which may stop a worker that reports it. A config cleanup runs + whatever is still recorded after a failure outside this wrapper. + + Args: + session: The pytest session. + + Returns: + The results of the other session-start hooks. + """ + config = session.config + # For a failure after this wrapper returns, such as in an outer wrapper. + config.add_cleanup(_run_cleanups) try: - _upload_data() - with contextlib.closing(connect()) as conn, conn.cursor() as cursor: - _create_database(cursor) - _create_tables(cursor) + if _owns_fixture_schema(config): + _cleanups.append(_drop_fixture_schema) + _create_fixture_schema() + if _is_test_process(config): + _cleanups.append(_delete_s3tables_namespace) + _create_s3tables_namespace() + _cleanups.append(functools.partial(_drop_database, ENV.schema)) + with contextlib.closing(connect()) as conn, conn.cursor() as cursor: + _create_database(cursor, ENV.schema) + return (yield) except BaseException: - _delete_s3tables_namespace() + # The original error is kept. + with contextlib.suppress(Exception): + _run_cleanups() raise def pytest_sessionfinish(session): - if not _is_test_process(session.config): - return - # Each cleanup step runs even if an earlier one fails. + """Check the fixture schema in the process that owns it, then remove what was created. + + The removal runs in this hook because a pytest-xdist worker reports that it + finished after it, and the controller may then stop the worker. + + Args: + session: The pytest session. + """ try: - with contextlib.closing(connect()) as conn, conn.cursor() as cursor: - _drop_database(cursor) + if _owns_fixture_schema(session.config): + _check_fixture_schema(session) finally: - try: - _delete_data() - finally: - _delete_s3tables_namespace() + _run_cleanups() + + +def _run_cleanups(): + """Run and forget the recorded removals, newest first, each even if one fails.""" + steps = list(reversed(_cleanups)) + _cleanups.clear() + _run_all(steps) def _is_test_process(config): @@ -56,6 +125,87 @@ def _is_test_process(config): return hasattr(config, "workerinput") or not getattr(config.option, "numprocesses", None) +def _owns_fixture_schema(config): + """Whether this process creates and drops the fixture schema. + + Args: + config: The pytest config. + + Returns: + True for the pytest-xdist controller or a run without workers, False for + a worker, which uses the controller's fixture schema. + """ + return not hasattr(config, "workerinput") + + +def _run_all(steps): + """Run every step in order, even after one fails, then raise the last failure. + + Args: + steps: Callables that take no arguments. + + Raises: + BaseException: The last exception a step raised, with the earlier ones + as its context. + """ + with contextlib.ExitStack() as stack: + for step in reversed(list(steps)): + stack.callback(step) + + +def _create_fixture_schema(): + """Upload the data files and create the fixture schema with its tables and views.""" + _upload_data() + with contextlib.closing(connect()) as conn, conn.cursor() as cursor: + _create_database(cursor, ENV.fixture_schema) + _create_tables(cursor) + + +def _drop_fixture_schema(): + """Drop the fixture schema and delete the uploaded data files.""" + _run_all([functools.partial(_drop_database, ENV.fixture_schema), _delete_data]) + + +def _check_fixture_schema(session): + """Fail the run if the fixture schema holds other tables than those it was created with. + + A missing fixture schema counts as holding no tables. Another failure to + list the tables is reported but does not fail the run. + + Args: + session: The pytest session, whose exit status is set to failed. + """ + expected = {t.name for t in TABLES} | {v.name for v in VIEWS} + try: + actual = { + table["Name"] + for page in boto3.client("glue") + .get_paginator("get_tables") + .paginate(DatabaseName=ENV.fixture_schema) + for table in page["TableList"] + } + except ClientError as e: + if e.response["Error"]["Code"] != "EntityNotFoundException": + sys.stderr.write( + f"\nCould not list the tables of fixture schema {ENV.fixture_schema}: {e!r}\n" + ) + return + actual = set() + except Exception as e: + sys.stderr.write( + f"\nCould not list the tables of fixture schema {ENV.fixture_schema}: {e!r}\n" + ) + return + if actual != expected: + sys.stderr.write( + f"\nFixture schema {ENV.fixture_schema} changed during the run; " + f"unexpected: {sorted(actual - expected)}, missing: {sorted(expected - actual)}. " + "Tests must create their objects in ENV.schema.\n" + ) + if session.exitstatus == pytest.ExitCode.OK: + session.exitstatus = pytest.ExitCode.TESTS_FAILED + + @functools.cache def _s3tables(): """Return an S3 Tables client and the ARN of ``ENV.s3tables_catalog``'s table bucket. @@ -90,17 +240,20 @@ def _create_s3tables_namespace(): def _delete_s3tables_namespace(): - """Delete this process's S3 Tables namespace and any table left in it.""" + """Delete this process's S3 Tables namespace and any table left in it, if it exists.""" if not ENV.s3tables_catalog: return client, arn = _s3tables() - tables = [ - table["name"] - for page in client.get_paginator("list_tables").paginate( - tableBucketARN=arn, namespace=ENV.s3tables_namespace - ) - for table in page["tables"] - ] + try: + tables = [ + table["name"] + for page in client.get_paginator("list_tables").paginate( + tableBucketARN=arn, namespace=ENV.s3tables_namespace + ) + for table in page["tables"] + ] + except client.exceptions.NotFoundException: + return for table in tables: client.delete_table(tableBucketARN=arn, namespace=ENV.s3tables_namespace, name=table) client.delete_namespace(tableBucketARN=arn, namespace=ENV.s3tables_namespace) @@ -108,12 +261,12 @@ def _delete_s3tables_namespace(): @functools.cache def _data_objects(): - """Return the S3 objects the session uploads: the table data files and test files. + """Return the S3 objects of the fixture schema: the table data files and test files. Returns: A dict from S3 key to object content. """ - prefix = f"{ENV.s3_staging_key}{ENV.schema}" + prefix = f"{ENV.s3_staging_key}{ENV.fixture_schema}" objects = { ENV.s3_filesystem_test_file_key: b"0123456789", f"{prefix}/spark_group_by/spark_group_by.csv": spark_group_by_csv(), @@ -139,27 +292,39 @@ def _delete_data(): client.delete_object(Bucket=ENV.s3_staging_bucket, Key=key) -def _create_database(cursor): - for q in read_query("create_database.sql.jinja2", schema=ENV.schema): +def _create_database(cursor, schema): + """Create a database. + + Args: + cursor: The cursor to run the statement with. + schema: The database name. + """ + for q in read_query("create_database.sql.jinja2", schema=schema): cursor.execute(q) -def _drop_database(cursor): - for q in read_query("drop_database.sql.jinja2", schema=ENV.schema): - cursor.execute(q) +def _drop_database(schema): + """Drop a database and its tables. + + Args: + schema: The database name. + """ + with contextlib.closing(connect()) as conn, conn.cursor() as cursor: + for q in read_query("drop_database.sql.jinja2", schema=schema): + cursor.execute(q) def _create_tables(cursor): - """Create the tables and views from ``tests.pyathena.tables``. + """Create the tables and views from ``tests.pyathena.tables`` in the fixture schema. Args: cursor: The cursor to run the statements with. """ for table in TABLES: - location = f"{ENV.s3_staging_dir}{ENV.schema}/{table.name}/" - cursor.execute(table.create_statement(ENV.schema, location)) + location = f"{ENV.s3_staging_dir}{ENV.fixture_schema}/{table.name}/" + cursor.execute(table.create_statement(ENV.fixture_schema, location)) for view in VIEWS: - cursor.execute(view.create_statement(ENV.schema)) + cursor.execute(view.create_statement(ENV.fixture_schema)) def connect(schema_name="default", **kwargs): @@ -171,6 +336,14 @@ def connect(schema_name="default", **kwargs): def create_engine(**kwargs): + """Create a SQLAlchemy engine whose default schema is the fixture schema. + + Args: + **kwargs: ``driver`` and the connection options to add to the URL. + + Returns: + The engine. + """ driver = kwargs.pop("driver", "rest") conn_str = SQLALCHEMY_CONNECTION_STRING.replace("+rest", f"+{driver}") for arg in [ @@ -197,7 +370,7 @@ def create_engine(**kwargs): return sqlalchemy.engine.create_engine( conn_str.format( region_name=ENV.region_name, - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, s3_staging_dir=ENV.s3_staging_dir, location=ENV.s3_staging_dir, **kwargs, @@ -206,6 +379,14 @@ def create_engine(**kwargs): def create_async_engine(**kwargs): + """Create an async SQLAlchemy engine whose default schema is the fixture schema. + + Args: + **kwargs: ``driver`` and the connection options to add to the URL. + + Returns: + The engine. + """ driver = kwargs.pop("driver", "aiorest") conn_str = ASYNC_SQLALCHEMY_CONNECTION_STRING.replace("+aiorest", f"+{driver}") if "unload" in kwargs: @@ -213,7 +394,7 @@ def create_async_engine(**kwargs): return _create_async_engine( conn_str.format( region_name=ENV.region_name, - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, s3_staging_dir=ENV.s3_staging_dir, location=ENV.s3_staging_dir, **kwargs, @@ -222,11 +403,20 @@ def create_async_engine(**kwargs): def _cursor(cursor_class, request): + """Yield a cursor whose default schema is the fixture schema. + + Args: + cursor_class: The cursor class. + request: The fixture request; its optional ``param`` holds connection options. + + Yields: + The cursor. + """ if not hasattr(request, "param"): request.param = {} with ( contextlib.closing( - connect(schema_name=ENV.schema, cursor_class=cursor_class, **request.param) + connect(schema_name=ENV.fixture_schema, cursor_class=cursor_class, **request.param) ) as conn, conn.cursor() as cursor, ): diff --git a/tests/pyathena/pandas/test_async_cursor.py b/tests/pyathena/pandas/test_async_cursor.py index a0e9261d..c37c30ea 100644 --- a/tests/pyathena/pandas/test_async_cursor.py +++ b/tests/pyathena/pandas/test_async_cursor.py @@ -225,7 +225,7 @@ def test_query_execution(self, async_pandas_cursor, parquet_engine, chunksize): future = async_pandas_cursor.query_execution(query_id) query_execution = future.result() - assert query_execution.database == ENV.schema + assert query_execution.database == ENV.fixture_schema assert query_execution.catalog assert query_execution.query_id if async_pandas_cursor._unload: diff --git a/tests/pyathena/pandas/test_util.py b/tests/pyathena/pandas/test_util.py index 1acbc07e..89ca76bb 100644 --- a/tests/pyathena/pandas/test_util.py +++ b/tests/pyathena/pandas/test_util.py @@ -377,7 +377,7 @@ def test_to_sql(cursor): compression="snappy", ) - cursor.execute(f"SELECT * FROM {table_name}") + cursor.execute(f"SELECT * FROM {ENV.schema}.{table_name}") assert cursor.fetchall() == [ ( 1, @@ -413,7 +413,7 @@ def test_to_sql(cursor): if_exists="append", compression="snappy", ) - cursor.execute(f"SELECT * FROM {table_name}") + cursor.execute(f"SELECT * FROM {ENV.schema}.{table_name}") assert cursor.fetchall() == [ ( 1, @@ -455,7 +455,7 @@ def test_to_sql_with_index(cursor): index=True, index_label="col_index", ) - cursor.execute(f"SELECT * FROM {table_name}") + cursor.execute(f"SELECT * FROM {ENV.schema}.{table_name}") assert cursor.fetchall() == [(0, 1)] assert [(d[0], d[1]) for d in cursor.description] == [ ("col_index", "bigint"), @@ -483,9 +483,9 @@ def test_to_sql_with_partitions(cursor): if_exists="fail", compression="snappy", ) - cursor.execute(f"SHOW PARTITIONS {table_name}") + cursor.execute(f"SHOW PARTITIONS {ENV.schema}.{table_name}") assert sorted(cursor.fetchall()) == [(f"col_int={i}",) for i in range(10)] - cursor.execute(f"SELECT COUNT(*) FROM {table_name}") + cursor.execute(f"SELECT COUNT(*) FROM {ENV.schema}.{table_name}") assert cursor.fetchall() == [(10,)] @@ -509,11 +509,11 @@ def test_to_sql_with_multiple_partitions(cursor): if_exists="fail", compression="snappy", ) - cursor.execute(f"SHOW PARTITIONS {table_name}") + cursor.execute(f"SHOW PARTITIONS {ENV.schema}.{table_name}") assert sorted(cursor.fetchall()), [(f"col_int={i}/col_string=a",) for i in range(5)] + [ (f"col_int={i}/col_string=b",) for i in range(5, 10) ] - cursor.execute(f"SELECT COUNT(*) FROM {table_name}") + cursor.execute(f"SELECT COUNT(*) FROM {ENV.schema}.{table_name}") assert cursor.fetchall() == [(10,)] diff --git a/tests/pyathena/polars/test_async_cursor.py b/tests/pyathena/polars/test_async_cursor.py index dde6c431..cd9198a4 100644 --- a/tests/pyathena/polars/test_async_cursor.py +++ b/tests/pyathena/polars/test_async_cursor.py @@ -121,7 +121,7 @@ def test_query_execution(self, async_polars_cursor): future = async_polars_cursor.query_execution(query_id) query_execution = future.result() - assert query_execution.database == ENV.schema + assert query_execution.database == ENV.fixture_schema assert query_execution.catalog assert query_execution.query_id if async_polars_cursor._unload: diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index c9da6169..11c24d26 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -455,7 +455,7 @@ def test_callback(query_id: str): def test_iter_chunks(self): """Test chunked iteration over query results.""" - with contextlib.closing(connect(schema_name=ENV.schema)) as conn: + with contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn: cursor = conn.cursor(PolarsCursor, chunksize=5) cursor.execute("SELECT * FROM many_rows LIMIT 15") chunks = list(cursor.iter_chunks()) @@ -476,7 +476,7 @@ def test_iter_chunks_without_chunksize(self, polars_cursor): def test_iter_chunks_many_rows(self): """Test chunked iteration with many rows.""" - with contextlib.closing(connect(schema_name=ENV.schema)) as conn: + with contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn: cursor = conn.cursor(PolarsCursor, chunksize=1000) cursor.execute("SELECT * FROM many_rows") chunks = list(cursor.iter_chunks()) @@ -505,7 +505,7 @@ def test_iter_chunks_unload(self, polars_cursor): def test_iter_chunks_data_consistency(self): """Test that chunked and regular reading produce the same data.""" - with contextlib.closing(connect(schema_name=ENV.schema)) as conn: + with contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn: # Regular reading (no chunksize) regular_cursor = conn.cursor(PolarsCursor) regular_cursor.execute("SELECT * FROM many_rows LIMIT 100") @@ -527,7 +527,7 @@ def test_iter_chunks_data_consistency(self): def test_iter_chunks_chunk_sizes(self): """Test that chunks have correct sizes.""" - with contextlib.closing(connect(schema_name=ENV.schema)) as conn: + with contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn: cursor = conn.cursor(PolarsCursor, chunksize=10) cursor.execute("SELECT * FROM many_rows LIMIT 50") @@ -550,7 +550,7 @@ def test_iter_chunks_chunk_sizes(self): def test_fetchone_with_chunksize(self): """Test that fetchone works correctly with chunksize enabled.""" - with contextlib.closing(connect(schema_name=ENV.schema)) as conn: + with contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn: cursor = conn.cursor(PolarsCursor, chunksize=5) cursor.execute("SELECT * FROM many_rows LIMIT 15") @@ -565,7 +565,7 @@ def test_fetchone_with_chunksize(self): def test_fetchmany_with_chunksize(self): """Test that fetchmany works correctly with chunksize enabled.""" - with contextlib.closing(connect(schema_name=ENV.schema)) as conn: + with contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn: cursor = conn.cursor(PolarsCursor, chunksize=5) cursor.execute("SELECT * FROM many_rows LIMIT 15") @@ -577,7 +577,7 @@ def test_fetchmany_with_chunksize(self): def test_fetchall_with_chunksize(self): """Test that fetchall works correctly with chunksize enabled.""" - with contextlib.closing(connect(schema_name=ENV.schema)) as conn: + with contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn: cursor = conn.cursor(PolarsCursor, chunksize=5) cursor.execute("SELECT * FROM many_rows LIMIT 15") @@ -586,7 +586,7 @@ def test_fetchall_with_chunksize(self): def test_iterator_with_chunksize(self): """Test that cursor iteration works correctly with chunksize enabled.""" - with contextlib.closing(connect(schema_name=ENV.schema)) as conn: + with contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn: cursor = conn.cursor(PolarsCursor, chunksize=5) cursor.execute("SELECT * FROM many_rows LIMIT 15") diff --git a/tests/pyathena/s3fs/test_async_cursor.py b/tests/pyathena/s3fs/test_async_cursor.py index ffc96dfd..2f151e41 100644 --- a/tests/pyathena/s3fs/test_async_cursor.py +++ b/tests/pyathena/s3fs/test_async_cursor.py @@ -151,7 +151,7 @@ def test_cancel(self, async_s3fs_cursor): def test_open_close(self): with ( - contextlib.closing(connect(schema_name=ENV.schema)) as conn, + contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn, conn.cursor(AsyncS3FSCursor) as cursor, ): query_id, future = cursor.execute("SELECT * FROM one_row") @@ -159,7 +159,7 @@ def test_open_close(self): assert result_set.fetchall() == [(1,)] def test_no_ops(self): - conn = connect(schema_name=ENV.schema) + conn = connect(schema_name=ENV.fixture_schema) cursor = conn.cursor(AsyncS3FSCursor) cursor.close() conn.close() diff --git a/tests/pyathena/s3fs/test_cursor.py b/tests/pyathena/s3fs/test_cursor.py index 40edaa07..dd04b7c3 100644 --- a/tests/pyathena/s3fs/test_cursor.py +++ b/tests/pyathena/s3fs/test_cursor.py @@ -148,14 +148,14 @@ def test_cancel_initial(self, s3fs_cursor): def test_open_close(self): with ( - contextlib.closing(connect(schema_name=ENV.schema)) as conn, + contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn, conn.cursor(S3FSCursor) as cursor, ): cursor.execute("SELECT * FROM one_row") assert cursor.fetchall() == [(1,)] def test_no_ops(self): - conn = connect(schema_name=ENV.schema) + conn = connect(schema_name=ENV.fixture_schema) cursor = conn.cursor(S3FSCursor) cursor.close() conn.close() @@ -247,7 +247,7 @@ def callback(query_id): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, cursor_kwargs={"on_start_query_execution": callback}, ) @@ -267,7 +267,7 @@ def callback(query_id): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, ) ) as conn, @@ -321,7 +321,7 @@ def test_basic_query_with_reader(self, csv_reader_class): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, cursor_kwargs={"csv_reader": csv_reader_class}, ) @@ -341,7 +341,7 @@ def test_multiple_columns_with_reader(self, csv_reader_class): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, cursor_kwargs={"csv_reader": csv_reader_class}, ) @@ -364,7 +364,7 @@ def test_null_with_default_reader(self): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, cursor_kwargs={"csv_reader": DefaultCSVReader}, ) @@ -380,7 +380,7 @@ def test_null_with_athena_reader(self): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, cursor_kwargs={"csv_reader": AthenaCSVReader}, ) @@ -396,7 +396,7 @@ def test_empty_string_with_default_reader(self): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, cursor_kwargs={"csv_reader": DefaultCSVReader}, ) @@ -413,7 +413,7 @@ def test_empty_string_with_athena_reader(self): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, cursor_kwargs={"csv_reader": AthenaCSVReader}, ) @@ -444,7 +444,7 @@ def test_null_vs_empty_string(self, csv_reader, expected_empty): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, cursor_kwargs={"csv_reader": csv_reader}, ) @@ -463,7 +463,7 @@ def test_mixed_values_with_athena_reader(self): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, cursor_kwargs={"csv_reader": AthenaCSVReader}, ) @@ -491,7 +491,7 @@ def test_quoted_string_with_comma(self, csv_reader_class): with ( contextlib.closing( connect( - schema_name=ENV.schema, + schema_name=ENV.fixture_schema, cursor_class=S3FSCursor, cursor_kwargs={"csv_reader": csv_reader_class}, ) diff --git a/tests/pyathena/spark/test_async_cursor.py b/tests/pyathena/spark/test_async_cursor.py index 97e72685..13987acb 100644 --- a/tests/pyathena/spark/test_async_cursor.py +++ b/tests/pyathena/spark/test_async_cursor.py @@ -45,7 +45,7 @@ def test_spark_dataframe(self, async_spark_cursor): df = spark.read.format("csv") \\ .option("header", "true") \\ .option("inferSchema", "true") \\ - .load("{ENV.s3_staging_dir}{ENV.schema}/spark_group_by/spark_group_by.csv") + .load("{ENV.s3_staging_dir}{ENV.fixture_schema}/spark_group_by/spark_group_by.csv") """ ), description="test description", @@ -95,7 +95,7 @@ def test_spark_sql(self, async_spark_cursor): query_id, future = async_spark_cursor.execute( textwrap.dedent( f""" - spark.sql("SELECT * FROM {ENV.schema}.one_row").show() + spark.sql("SELECT * FROM {ENV.fixture_schema}.one_row").show() """ ) ) diff --git a/tests/pyathena/spark/test_spark_cursor.py b/tests/pyathena/spark/test_spark_cursor.py index 15267854..96cd7e34 100644 --- a/tests/pyathena/spark/test_spark_cursor.py +++ b/tests/pyathena/spark/test_spark_cursor.py @@ -31,7 +31,7 @@ def test_spark_dataframe(self, spark_cursor): df = spark.read.format("csv") \\ .option("header", "true") \\ .option("inferSchema", "true") \\ - .load("{ENV.s3_staging_dir}{ENV.schema}/spark_group_by/spark_group_by.csv") + .load("{ENV.s3_staging_dir}{ENV.fixture_schema}/spark_group_by/spark_group_by.csv") """ ), description="test description", @@ -91,7 +91,7 @@ def test_spark_sql(self, spark_cursor): spark_cursor.execute( textwrap.dedent( f""" - spark.sql("SELECT * FROM {ENV.schema}.one_row").show() + spark.sql("SELECT * FROM {ENV.fixture_schema}.one_row").show() """ ) ) diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 69e91352..f4f2954b 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -985,7 +985,7 @@ def test_reflect_no_such_table(self, engine): def test_reflect_table(self, engine): engine, conn = engine - one_row = Table("one_row", MetaData(schema=ENV.schema), autoload_with=conn) + one_row = Table("one_row", MetaData(schema=ENV.fixture_schema), autoload_with=conn) assert len(one_row.c) == 1 assert one_row.c.number_of_rows is not None assert one_row.comment == "table comment" @@ -996,7 +996,7 @@ def test_reflect_table(self, engine): assert "file_format" in dialect_opts assert "serdeproperties" in dialect_opts assert "tblproperties" in dialect_opts - assert dialect_opts["location"] == f"{ENV.s3_staging_dir}{ENV.schema}/one_row" + assert dialect_opts["location"] == f"{ENV.s3_staging_dir}{ENV.fixture_schema}/one_row" assert ( dialect_opts["row_format"] == "SERDE 'org.apache.hadoop.hive.serde2.lazy.LazySimpleSerDe'" @@ -1014,7 +1014,7 @@ def test_reflect_table(self, engine): def test_reflect_table_with_schema(self, engine): engine, conn = engine - one_row = Table("one_row", MetaData(schema=ENV.schema), autoload_with=conn) + one_row = Table("one_row", MetaData(schema=ENV.fixture_schema), autoload_with=conn) assert len(one_row.c) == 1 assert one_row.c.number_of_rows is not None assert one_row.comment == "table comment" @@ -1025,7 +1025,7 @@ def test_reflect_table_with_schema(self, engine): assert "file_format" in dialect_opts assert "serdeproperties" in dialect_opts assert "tblproperties" in dialect_opts - assert dialect_opts["location"] == f"{ENV.s3_staging_dir}{ENV.schema}/one_row" + assert dialect_opts["location"] == f"{ENV.s3_staging_dir}{ENV.fixture_schema}/one_row" assert ( dialect_opts["row_format"] == "SERDE 'org.apache.hadoop.hive.serde2.lazy.LazySimpleSerDe'" @@ -1043,7 +1043,7 @@ def test_reflect_table_with_schema(self, engine): def test_reflect_table_include_columns(self, engine): engine, conn = engine - one_row_complex = Table("one_row_complex", MetaData(schema=ENV.schema)) + one_row_complex = Table("one_row_complex", MetaData(schema=ENV.fixture_schema)) insp = sqlalchemy.inspect(engine) insp.reflect_table( one_row_complex, @@ -1057,7 +1057,9 @@ def test_reflect_table_include_columns(self, engine): def test_partition_table_columns(self, engine): engine, conn = engine - partition_table = Table("partition_table", MetaData(schema=ENV.schema), autoload_with=conn) + partition_table = Table( + "partition_table", MetaData(schema=ENV.fixture_schema), autoload_with=conn + ) assert len(partition_table.columns) == 2 assert "a" in partition_table.columns assert "b" in partition_table.columns @@ -1074,38 +1076,38 @@ def test_reflect_schemas(self, engine): engine, conn = engine insp = sqlalchemy.inspect(engine) schemas = insp.get_schema_names() - assert ENV.schema in schemas + assert ENV.fixture_schema in schemas assert "default" in schemas def test_get_table_names(self, engine): engine, conn = engine - meta = MetaData(schema=ENV.schema) + meta = MetaData(schema=ENV.fixture_schema) meta.reflect(bind=engine) # With schema specified, table names are schema-qualified - schema_qualified_one_row = f"{ENV.schema}.one_row" - schema_qualified_one_row_complex = f"{ENV.schema}.one_row_complex" - schema_qualified_view_one_row = f"{ENV.schema}.view_one_row" + schema_qualified_one_row = f"{ENV.fixture_schema}.one_row" + schema_qualified_one_row_complex = f"{ENV.fixture_schema}.one_row_complex" + schema_qualified_view_one_row = f"{ENV.fixture_schema}.view_one_row" assert schema_qualified_one_row in meta.tables assert schema_qualified_one_row_complex in meta.tables assert schema_qualified_view_one_row not in meta.tables insp = sqlalchemy.inspect(engine) - assert "many_rows" in insp.get_table_names(schema=ENV.schema) + assert "many_rows" in insp.get_table_names(schema=ENV.fixture_schema) def test_get_view_names(self, engine): engine, conn = engine - meta = MetaData(schema=ENV.schema) + meta = MetaData(schema=ENV.fixture_schema) meta.reflect(bind=engine, views=True) # With schema specified, table names are schema-qualified - schema_qualified_one_row = f"{ENV.schema}.one_row" - schema_qualified_one_row_complex = f"{ENV.schema}.one_row_complex" - schema_qualified_view_one_row = f"{ENV.schema}.view_one_row" + schema_qualified_one_row = f"{ENV.fixture_schema}.one_row" + schema_qualified_one_row_complex = f"{ENV.fixture_schema}.one_row_complex" + schema_qualified_view_one_row = f"{ENV.fixture_schema}.view_one_row" assert schema_qualified_one_row in meta.tables assert schema_qualified_one_row_complex in meta.tables assert schema_qualified_view_one_row in meta.tables insp = sqlalchemy.inspect(engine) - actual = insp.get_view_names(schema=ENV.schema) + actual = insp.get_view_names(schema=ENV.fixture_schema) assert "one_row" not in actual assert "one_row_complex" not in actual assert "view_one_row" in actual @@ -1113,16 +1115,16 @@ def test_get_view_names(self, engine): def test_get_table_comment(self, engine): engine, conn = engine insp = sqlalchemy.inspect(engine) - actual = insp.get_table_comment("one_row", schema=ENV.schema) + actual = insp.get_table_comment("one_row", schema=ENV.fixture_schema) assert actual == {"text": "table comment"} def test_get_table_options(self, engine): engine, conn = engine insp = sqlalchemy.inspect(engine) - actual = insp.get_table_options("parquet_with_compression", schema=ENV.schema) + actual = insp.get_table_options("parquet_with_compression", schema=ENV.fixture_schema) assert ( actual["awsathena_location"] - == f"{ENV.s3_staging_dir}{ENV.schema}/parquet_with_compression" + == f"{ENV.s3_staging_dir}{ENV.fixture_schema}/parquet_with_compression" ) assert actual["awsathena_compression"] == "SNAPPY" assert ( @@ -1142,13 +1144,13 @@ def test_get_table_options(self, engine): def test_has_table(self, engine): engine, conn = engine insp = sqlalchemy.inspect(engine) - assert insp.has_table("one_row", schema=ENV.schema) - assert not insp.has_table("this_table_does_not_exist", schema=ENV.schema) + assert insp.has_table("one_row", schema=ENV.fixture_schema) + assert not insp.has_table("this_table_does_not_exist", schema=ENV.fixture_schema) def test_get_columns(self, engine): engine, conn = engine insp = sqlalchemy.inspect(engine) - actual = insp.get_columns(table_name="one_row", schema=ENV.schema)[0] + actual = insp.get_columns(table_name="one_row", schema=ENV.fixture_schema)[0] assert actual["name"] == "number_of_rows" assert isinstance(actual["type"], types.INTEGER) assert actual["nullable"] @@ -1169,16 +1171,16 @@ def reflect(): return ( { t: ( - insp.get_table_comment(t, schema=ENV.schema), - insp.get_table_options(t, schema=ENV.schema), + insp.get_table_comment(t, schema=ENV.fixture_schema), + insp.get_table_options(t, schema=ENV.fixture_schema), ) for t in tables }, - insp.get_table_names(schema=ENV.schema), - insp.get_view_names(schema=ENV.schema), + insp.get_table_names(schema=ENV.fixture_schema), + insp.get_view_names(schema=ENV.fixture_schema), # Other runs create and drop schemas concurrently, so only this # run's schema is compared. - ENV.schema in insp.get_schema_names(), + ENV.fixture_schema in insp.get_schema_names(), ) expected = reflect() @@ -1354,7 +1356,9 @@ def test_get_view_definition_across_cursor_types(self, engine): def test_char_length(self, engine): engine, conn = engine - one_row_complex = Table("one_row_complex", MetaData(schema=ENV.schema), autoload_with=conn) + one_row_complex = Table( + "one_row_complex", MetaData(schema=ENV.fixture_schema), autoload_with=conn + ) result = conn.execute( sqlalchemy.select(sqlalchemy.func.char_length(one_row_complex.c.col_string)) ).scalar() @@ -1362,7 +1366,9 @@ def test_char_length(self, engine): def test_filter_func(self, engine): engine, conn = engine - one_row_complex = Table("one_row_complex", MetaData(schema=ENV.schema), autoload_with=conn) + one_row_complex = Table( + "one_row_complex", MetaData(schema=ENV.fixture_schema), autoload_with=conn + ) # Test filter() function basic functionality # @@ -1421,7 +1427,9 @@ def test_filter_func(self, engine): def test_reflect_select(self, engine): engine, conn = engine - one_row_complex = Table("one_row_complex", MetaData(schema=ENV.schema), autoload_with=conn) + one_row_complex = Table( + "one_row_complex", MetaData(schema=ENV.fixture_schema), autoload_with=conn + ) assert len(one_row_complex.c) == 16 assert isinstance(one_row_complex.c.col_string, Column) rows = conn.execute(one_row_complex.select()).fetchall() @@ -1472,7 +1480,7 @@ def test_reflect_select(self, engine): def test_select_offset_limit(self, engine): engine, conn = engine - many_rows = Table("many_rows", MetaData(schema=ENV.schema), autoload_with=conn) + many_rows = Table("many_rows", MetaData(schema=ENV.fixture_schema), autoload_with=conn) rows = conn.execute(many_rows.select().offset(10).limit(5)).fetchall() assert rows == [(i,) for i in range(10, 15)] @@ -2426,7 +2434,7 @@ def test_cast_as_varchar(self, engine): engine, conn = engine # varchar without length - one_row = Table("one_row", MetaData(schema=ENV.schema), autoload_with=conn) + one_row = Table("one_row", MetaData(schema=ENV.fixture_schema), autoload_with=conn) actual = conn.execute( sqlalchemy.select(expression.cast(one_row.c.number_of_rows, types.VARCHAR)) ).scalar() @@ -2499,7 +2507,9 @@ def test_binary_null_vs_empty(self, engine): def test_cast_as_binary(self, engine): engine, conn = engine - one_row_complex = Table("one_row_complex", MetaData(schema=ENV.schema), autoload_with=conn) + one_row_complex = Table( + "one_row_complex", MetaData(schema=ENV.fixture_schema), autoload_with=conn + ) actual = conn.execute( sqlalchemy.select( expression.cast(one_row_complex.c.col_string, types.BINARY), @@ -3279,15 +3289,15 @@ def record(conn, cursor, statement, parameters, context, executemany): def test_get_view_definition(self, engine): engine, conn = engine insp = sqlalchemy.inspect(engine) - actual = insp.get_view_definition(schema=ENV.schema, view_name="v_one_row") + actual = insp.get_view_definition(schema=ENV.fixture_schema, view_name="v_one_row") assert ( actual == textwrap.dedent( f""" - CREATE VIEW {ENV.schema}.v_one_row AS + CREATE VIEW {ENV.fixture_schema}.v_one_row AS SELECT number_of_rows FROM - {ENV.schema}.one_row + {ENV.fixture_schema}.one_row """ ).strip() ) @@ -3297,7 +3307,7 @@ def test_get_view_definition_missing_view(self, engine): insp = sqlalchemy.inspect(engine) pytest.raises( NoSuchTableError, - lambda: insp.get_view_definition(schema=ENV.schema, view_name="test_view"), + lambda: insp.get_view_definition(schema=ENV.fixture_schema, view_name="test_view"), ) def test_numeric_type_variants(self, engine): diff --git a/tests/pyathena/test_async_cursor.py b/tests/pyathena/test_async_cursor.py index 1cf4d0d9..148436e6 100644 --- a/tests/pyathena/test_async_cursor.py +++ b/tests/pyathena/test_async_cursor.py @@ -21,7 +21,7 @@ def test_fetchone(self, async_cursor): assert result_set.fetchone() == (1,) assert result_set.rownumber == 1 assert result_set.fetchone() is None - assert result_set.database == ENV.schema + assert result_set.database == ENV.fixture_schema assert result_set.catalog assert result_set.query_id assert result_set.query @@ -101,7 +101,7 @@ def test_query_execution(self, async_cursor): future = async_cursor.query_execution(query_id) query_execution = future.result() - assert query_execution.database == ENV.schema + assert query_execution.database == ENV.fixture_schema assert query_execution.catalog assert query_execution.query_id assert query_execution.query == query diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index aa1832b3..c93e84fc 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -49,7 +49,7 @@ def test_fetchone(self, cursor): assert cursor.fetchone() == (1,) assert cursor.rownumber == 1 assert cursor.fetchone() is None - assert cursor.database == ENV.schema + assert cursor.database == ENV.fixture_schema assert cursor.catalog assert cursor.query_id assert cursor.query @@ -797,7 +797,7 @@ def test_cancel_initial(self, cursor): def test_multiple_connection(self): def execute_other_thread(): with ( - contextlib.closing(connect(schema_name=ENV.schema)) as conn, + contextlib.closing(connect(schema_name=ENV.fixture_schema)) as conn, conn.cursor() as cursor, ): cursor.execute("SELECT * FROM one_row") @@ -1397,7 +1397,7 @@ def read(): for t in ("one_row", "parquet_with_compression", "partition_table") ], sorted(self._metadata_view(m) for m in cursor.list_table_metadata()), - ENV.schema in [d.name for d in cursor.list_databases("AwsDataCatalog")], + ENV.fixture_schema in [d.name for d in cursor.list_databases("AwsDataCatalog")], ) expected = read() diff --git a/tests/pyathena/test_glue.py b/tests/pyathena/test_glue.py index 82d786fa..11ec70cf 100644 --- a/tests/pyathena/test_glue.py +++ b/tests/pyathena/test_glue.py @@ -48,14 +48,16 @@ def test_reads_what_athena_reports(self, cursor): catalog = cursor.connection.catalog_name for table in ("one_row", "parquet_with_compression", "partition_table", "view_one_row"): - assert self._view(glue.get_table(catalog, ENV.schema, table)) == self._view( + assert self._view(glue.get_table(catalog, ENV.fixture_schema, table)) == self._view( cursor.get_table_metadata(table) ) - assert sorted(self._view(m) for m in glue.list_tables(catalog, ENV.schema)) == sorted( - self._view(m) for m in cursor.list_table_metadata() - ) - assert [m.name for m in glue.list_tables(catalog, ENV.schema, "one_row")] == ["one_row"] - assert ENV.schema in [d.name for d in glue.list_databases(catalog)] + assert sorted( + self._view(m) for m in glue.list_tables(catalog, ENV.fixture_schema) + ) == sorted(self._view(m) for m in cursor.list_table_metadata()) + assert [m.name for m in glue.list_tables(catalog, ENV.fixture_schema, "one_row")] == [ + "one_row" + ] + assert ENV.fixture_schema in [d.name for d in glue.list_databases(catalog)] @pytest.mark.skipif( not ENV.s3tables_catalog, @@ -94,7 +96,7 @@ def test_reads_s3_tables_catalog(self): def test_reports_a_missing_table(self, cursor): with pytest.raises(ClientError) as caught: cursor.connection._glue.get_table( - cursor.connection.catalog_name, ENV.schema, "no_such_table_786" + cursor.connection.catalog_name, ENV.fixture_schema, "no_such_table_786" ) assert caught.value.response["Error"]["Code"] == "EntityNotFoundException" @@ -102,13 +104,13 @@ def test_reports_a_missing_table(self, cursor): def test_rejects_an_unsupported_catalog(self, cursor): with pytest.raises(ValueError, match="federated_catalog"): - cursor.connection._glue.get_table("federated_catalog", ENV.schema, "one_row") + cursor.connection._glue.get_table("federated_catalog", ENV.fixture_schema, "one_row") def test_stops_after_it_cannot_reach_glue(self, cursor): glue = unreachable_glue(cursor.connection) with pytest.raises(BotoConnectionError): - glue.get_table("AwsDataCatalog", ENV.schema, "one_row") + glue.get_table("AwsDataCatalog", ENV.fixture_schema, "one_row") assert not glue.reachable assert not glue.usable_for("AwsDataCatalog")