Skip to content
Merged
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
1 change: 0 additions & 1 deletion scripts/config/license_headers.toml
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,5 @@ unheaded-files = [
"tests/pyathena/test_cursor.py",
"tests/pyathena/test_model.py",
"tests/pyathena/test_util.py",
"tests/resources/queries/create_table.sql.jinja2",
"tests/sqlalchemy/test_suite.py",
]
66 changes: 41 additions & 25 deletions tests/pyathena/conftest.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,14 @@
import contextlib
import functools
import uuid
from io import BytesIO
from pathlib import Path

import boto3
import pytest
import sqlalchemy
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


Expand All @@ -21,10 +20,10 @@ def pytest_sessionstart(session):
# pytest skips pytest_sessionfinish after a failed pytest_sessionstart, so
# a failure after the namespace is created deletes it here.
try:
_upload_rows()
_upload_data()
with contextlib.closing(connect()) as conn, conn.cursor() as cursor:
_create_database(cursor)
_create_table(cursor)
_create_tables(cursor)
except BaseException:
_delete_s3tables_namespace()
raise
Expand All @@ -39,7 +38,7 @@ def pytest_sessionfinish(session):
_drop_database(cursor)
finally:
try:
_delete_rows()
_delete_data()
finally:
_delete_s3tables_namespace()

Expand Down Expand Up @@ -107,26 +106,37 @@ def _delete_s3tables_namespace():
client.delete_namespace(tableBucketARN=arn, namespace=ENV.s3tables_namespace)


def _upload_rows():
@functools.cache
def _data_objects():

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Independent review (relayed Codex result): CLEAN

  • Reviewer: Codex CLI 0.157.1 (codex exec -s read-only), model gpt-6-astra, session 01a0e72e-c8d6-7830-88e2-bb7c702d0fd3.
  • Scope: base 7b4cceb, head 2760077. The review ran on a detached snapshot worktree without .env. The prompt left out the PR number, description, commit messages, and self-review findings. After the review, the snapshot was unchanged and still at the reviewed head.
  • Covered surfaces, as reported:
    • the generated DDL and data compared with all removed resources, including the decompressed Hive row;
    • types, values, NULLs, row order, reflection metadata, views, and the consumers of the Spark CSV;
    • S3 keys, upload/delete symmetry, worker isolation, and setup/cleanup failure paths;
    • Python and PyArrow compatibility, workflow dependencies, and remaining references to removed files.
  • Result: "No correctness regressions found by static inspection." This was a static review only: no builds, tests, or network access.

"""Return the S3 objects the session uploads: the table data files and test files.

Returns:
A dict from S3 key to object content.
"""
prefix = f"{ENV.s3_staging_key}{ENV.schema}"
objects = {
ENV.s3_filesystem_test_file_key: b"0123456789",
f"{prefix}/spark_group_by/spark_group_by.csv": spark_group_by_csv(),
}
for table in TABLES:
if data_file := table.data_file():
name, content = data_file
objects[f"{prefix}/{table.name}/{name}"] = content
return objects


def _upload_data():
"""Upload the objects from ``_data_objects``."""
client = boto3.client("s3")
rows = Path(__file__).parents[1].resolve() / "resources" / "rows"
for row in rows.iterdir():
key = f"{ENV.s3_staging_key}{ENV.schema}/{row.stem}/{row.name}"
client.upload_file(str(row), ENV.s3_staging_bucket, key)
client.upload_fileobj(
BytesIO(b"0123456789"),
ENV.s3_staging_bucket,
ENV.s3_filesystem_test_file_key,
)
for key, content in _data_objects().items():
client.put_object(Bucket=ENV.s3_staging_bucket, Key=key, Body=content)


def _delete_rows():
def _delete_data():
"""Delete the objects from ``_data_objects``."""
client = boto3.client("s3")
rows = Path(__file__).parents[1].resolve() / "resources" / "rows"
for row in rows.iterdir():
key = f"{ENV.s3_staging_key}{ENV.schema}/{row.stem}/{row.name}"
for key in _data_objects():
client.delete_object(Bucket=ENV.s3_staging_bucket, Key=key)
client.delete_object(Bucket=ENV.s3_staging_bucket, Key=ENV.s3_filesystem_test_file_key)


def _create_database(cursor):
Expand All @@ -139,11 +149,17 @@ def _drop_database(cursor):
cursor.execute(q)


def _create_table(cursor):
for q in read_query(
"create_table.sql.jinja2", s3_staging_dir=ENV.s3_staging_dir, schema=ENV.schema
):
cursor.execute(q)
def _create_tables(cursor):
"""Create the tables and views from ``tests.pyathena.tables``.

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))
for view in VIEWS:
cursor.execute(view.create_statement(ENV.schema))


def connect(schema_name="default", **kwargs):
Expand Down
260 changes: 260 additions & 0 deletions tests/pyathena/tables.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,260 @@
# Copyright 2026 The PyAthena authors
#
# Licensed under the MIT License.
# See LICENSE or https://opensource.org/licenses/MIT.
#
# SPDX-License-Identifier: MIT
"""The shared tables and views the PyAthena test session creates.

Each table is defined once here. The session setup generates its DDL and its
data file from the definition.
"""

import csv
import io
from dataclasses import dataclass
from datetime import date, datetime
from decimal import Decimal
from typing import Any, Literal

import pyarrow as pa
import pyarrow.parquet as pq

_TEXT_FORMAT = (
"ROW FORMAT DELIMITED FIELDS TERMINATED BY '\\t' LINES TERMINATED BY '\\n' STORED AS TEXTFILE"
)


@dataclass(frozen=True)
class Column:
"""A table column.

Attributes:
name: The column name.
athena_type: The column type in Athena DDL.
arrow_type: The Arrow type of the column in a Parquet data file.
comment: The column comment.
"""

name: str
athena_type: str
arrow_type: pa.DataType
comment: str | None = None

def ddl(self) -> str:
"""Return the column definition for ``CREATE TABLE``.

Returns:
The column name, type, and comment.
"""
comment = f" COMMENT '{self.comment}'" if self.comment else ""
return f"{self.name} {self.athena_type}{comment}"


@dataclass(frozen=True)
class Table:
"""A table with its data.

Attributes:
name: The table name.
columns: The data columns.
rows: The rows, as tuples of Python values in column order.
storage: ``"parquet"`` or ``"text"`` (tab-separated values).
partitions: The partition columns. The session adds no partitions.
comment: The table comment.
tblproperties: The table properties.
"""

name: str
columns: tuple[Column, ...]
rows: tuple[tuple[Any, ...], ...] = ()
storage: Literal["parquet", "text"] = "parquet"

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Self-review round one (implementation behavior): FINDINGS (1 simplification, fixed)

Base 7b4cceb (#873 head, the stacked base), head 12cfd65. Full diff reviewed (10 files).

  • Finding: Table.storage was a plain str. A misspelled value would silently produce Parquet DDL and data. It is now Literal["parquet", "text"], in 2760077. No behavior change; just lint and mypy tests/pyathena/tables.py pass.
  • DDL:
    • Text tables emit the old ROW FORMAT DELIMITED ... STORED AS TEXTFILE clause verbatim, so the one_row reflection asserts (SerDe, input/output format, field.delim, line.delim, serialization.format) still hold.
    • parquet_with_compression keeps STORED AS PARQUET and TBLPROPERTIES ('parquet.compress'='SNAPPY').
    • CREATE VIEW replaces CREATE OR REPLACE VIEW, which is safe because every session uses a fresh schema.
  • Data:
    • Parquet is written with an explicit schema from Column.arrow_type. The live Athena results match the old files: see the parity check in the PR description.
    • timestamp("ms"), map_ from key/value tuples, decimal128(10,1) and binary all read back as before, and test_complex passes for every cursor.
    • NULLs in integer_na_values/boolean_na_values are Parquet nulls, and the NA tests pass.
  • Row order: many_rows stays text after the observed Parquet reordering. The 3-row NA tables are single-page Parquet files, and their SELECT * order tests pass.
  • Resources: _data_objects is cached per process (ENV.schema is fixed at import), so setup uploads and teardown deletes the same keys. When the upload fails midway, pytest_sessionfinish is skipped, as before this change, and the 1-day lifecycle rule removes the objects.
  • Python 3.10: str | None annotations and zip(strict=True) are 3.10-compatible. pa.Table.from_pylist exists in the pyarrow>=10 floor.

partitions: tuple[Column, ...] = ()
comment: str | None = None
tblproperties: tuple[tuple[str, str], ...] = ()

def create_statement(self, schema: str, location: str) -> str:
"""Return the ``CREATE EXTERNAL TABLE`` statement.

Args:
schema: The schema to create the table in.
location: The S3 location of the table data, ending with a slash.

Returns:
The statement.
"""
columns = ",\n ".join(c.ddl() for c in self.columns)
clauses = [f"CREATE EXTERNAL TABLE {schema}.{self.name} (\n {columns}\n)"]
if self.comment:
clauses.append(f"COMMENT '{self.comment}'")
if self.partitions:
clauses.append(f"PARTITIONED BY ({', '.join(c.ddl() for c in self.partitions)})")
clauses.append(_TEXT_FORMAT if self.storage == "text" else "STORED AS PARQUET")
clauses.append(f"LOCATION '{location}'")
if self.tblproperties:
properties = ", ".join(f"'{k}'='{v}'" for k, v in self.tblproperties)
clauses.append(f"TBLPROPERTIES ({properties})")
return "\n".join(clauses)

def data_file(self) -> tuple[str, bytes] | None:
"""Return the table's data file.

Returns:
The file name and content, or None for a table without rows.
"""
if not self.rows:
return None
if self.storage == "text":
lines = ["\t".join(_to_text(v) for v in row) for row in self.rows]
return "data.tsv", "".join(f"{line}\n" for line in lines).encode()
schema = pa.schema([(c.name, c.arrow_type) for c in self.columns])
table = pa.Table.from_pylist(
[dict(zip(schema.names, row, strict=True)) for row in self.rows], schema=schema
)
buffer = io.BytesIO()
pq.write_table(table, buffer)
return "data.parquet", buffer.getvalue()


@dataclass(frozen=True)
class View:
"""A view.

Attributes:
name: The view name.
query: The view query; ``{schema}`` is replaced with the schema name.
"""

name: str
query: str

def create_statement(self, schema: str) -> str:
"""Return the ``CREATE VIEW`` statement.

Args:
schema: The schema to create the view in.

Returns:
The statement.
"""
return f"CREATE VIEW {schema}.{self.name} AS {self.query.format(schema=schema)}"


def _to_text(value: Any) -> str:
"""Format a scalar value for a tab-separated text file.

Args:
value: The value.

Returns:
The text; an empty string for None.
"""
if value is None:
return ""
if isinstance(value, bool):
return str(value).lower()
return str(value)


TABLES = (
# A text table: the reflection tests assert its SerDe and delimiters.
Table(
"one_row",
(Column("number_of_rows", "INT", pa.int32(), comment="some comment"),),
rows=((1,),),
storage="text",
comment="table comment",
),
# A text table: tests read it without ORDER BY and expect the file order,

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Self-review round two (claims, callers, AWS operations): FINDINGS (description only, fixed)

Base 7b4cceb, head 2760077.

Claims checked:

  • "Same statements as before": after Give the executemany and partition tests their own tables聽#873 the template had 7 tables and 2 views, and TABLES/VIEWS have the same 7 and 2.
  • "7 objects per process": before, 6 row files plus test.dat; now 5 data files (partition_table and parquet_with_compression have no rows) plus the Spark CSV plus test.dat.
  • "Byte-identical": checked offline for the one_row, many_rows and spark_group_by.csv bytes.
  • "Out of file order": seen in the first targeted run, where test_pandas_cursor_chunked_vs_regular_same_data got 0..6, 15..30, 7... After the switch to text, the order-sensitive -k set (371 passed) includes both failed tests.
  • Removed-file references: git grep over the repository finds no remaining reference to tests/resources/rows, create_table.sql or the row file names. The license-header config entry is removed, and NOTICE and docs never listed them.
  • CI path filters: nothing under pyathena/(aio/)?sqlalchemy, tests/sqlalchemy, tests/pyathena/(aio/)?sqlalchemy or the Spark paths changed, so CI runs only the PyAthena suite, without Spark tests.

Finding: the TEST section attributed the Spark test_spark_sql skips to xdist without evidence. The skip reason printed by -rs is pytest-dependency's depends on test_spark_dataframe. The description now quotes that reason and notes the Literal-only commit after the tested commit.

# which Athena does not keep for a Parquet file of this size.
Table(
"many_rows",
(Column("a", "INT", pa.int32()),),
rows=tuple((i,) for i in range(10000)),
storage="text",
),
Table(
"one_row_complex",
(
Column("col_boolean", "BOOLEAN", pa.bool_()),
Column("col_tinyint", "TINYINT", pa.int8()),
Column("col_smallint", "SMALLINT", pa.int16()),
Column("col_int", "INT", pa.int32()),
Column("col_bigint", "BIGINT", pa.int64()),
Column("col_float", "FLOAT", pa.float32()),
Column("col_double", "DOUBLE", pa.float64()),
Column("col_string", "STRING", pa.string()),
Column("col_varchar", "VARCHAR(10)", pa.string()),
Column("col_timestamp", "TIMESTAMP", pa.timestamp("ms")),
Column("col_date", "DATE", pa.date32()),
Column("col_binary", "BINARY", pa.binary()),
Column("col_array", "ARRAY<int>", pa.list_(pa.int32())),
Column("col_map", "MAP<int, int>", pa.map_(pa.int32(), pa.int32())),
Column(
"col_struct",
"STRUCT<a: int, b: int>",
pa.struct([("a", pa.int32()), ("b", pa.int32())]),
),
Column("col_decimal", "DECIMAL(10,1)", pa.decimal128(10, 1)),
),
rows=(
(
True,
127,
32767,
2147483647,
9223372036854775807,
0.5,
0.25,
"a string",
"varchar",
datetime(2017, 1, 1, 0, 0, 0),
date(2017, 1, 2),
b"123",
[1, 2],
[(1, 2), (3, 4)],
{"a": 1, "b": 2},
Decimal("0.1"),
),
),
),
Table(
"partition_table",
(Column("a", "STRING", pa.string()),),
partitions=(Column("b", "INT", pa.int32()),),
),
Table(
"integer_na_values",
(Column("a", "INT", pa.int32()), Column("b", "INT", pa.int32())),
rows=((1, 2), (1, None), (None, None)),
),
Table(
"boolean_na_values",
(Column("a", "BOOLEAN", pa.bool_()), Column("b", "BOOLEAN", pa.bool_())),
rows=((True, False), (False, None), (None, None)),
),
Table(
"parquet_with_compression",
(Column("a", "INT", pa.int32()),),
tblproperties=(("parquet.compress", "SNAPPY"),),
),
)

VIEWS = (
View("view_one_row", "SELECT * FROM {schema}.one_row"),
View("v_one_row", "SELECT number_of_rows FROM {schema}.one_row"),
)

# The Spark tests read this CSV file from ``<schema>/spark_group_by/``.
SPARK_GROUP_BY = (("name", "count"), ("foo", 1), ("bar", 2), ("bar", 3), ("foo", 4))


def spark_group_by_csv() -> bytes:
"""Return ``SPARK_GROUP_BY`` as a CSV file with a header row.

Returns:
The file content.
"""
buffer = io.StringIO()
csv.writer(buffer, lineterminator="\n").writerows(SPARK_GROUP_BY)
return buffer.getvalue().encode()
Loading
Loading