From d6b930356837a7b39c8ed99bf4548e2dc9925e39 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Wed, 30 Sep 2026 15:02:59 +0900 Subject: [PATCH] Align the compliance-suite dialect registry with the entry points tests/sqlalchemy/__init__.py registers the dialects for the SQLAlchemy compliance suite, and a registration overrides the installed entry point in any process that imports it. After #839 pointed the bare awsathena entry point at AthenaRestDialect, the file still registered the base AthenaDialect, and it never registered awsathena.polars. Register the same classes as pyproject.toml, and add a test that compares the file's registrations with the installed entry points. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/sqlalchemy/test_base.py | 18 ++++++++++++++++++ tests/sqlalchemy/__init__.py | 3 ++- 2 files changed, 20 insertions(+), 1 deletion(-) diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 273681ebb..8af87da47 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -1,9 +1,12 @@ import contextlib +import importlib.metadata import re +import runpy import textwrap import uuid from datetime import date, datetime from decimal import Decimal +from pathlib import Path from types import SimpleNamespace from urllib.parse import quote_plus @@ -18,6 +21,7 @@ from sqlalchemy.sql.ddl import CreateTable from sqlalchemy.sql.schema import Column, MetaData, Table from sqlalchemy.sql.selectable import TextualSelect +from sqlalchemy.util import PluginLoader from pyathena.aio.sqlalchemy.base import AthenaAioDialect from pyathena.converter import DefaultTypeConverter @@ -122,6 +126,20 @@ def test_bare_scheme_uses_rest_driver(self): assert bare.url.get_driver_name() == "rest" assert bare.dialect.dialect_description == "awsathena+rest" + def test_compliance_suite_registry_matches_entry_points(self, monkeypatch): + # tests/sqlalchemy registers the dialects for the compliance suite, and a + # registration overrides the installed entry point in that process. + loader = PluginLoader("sqlalchemy.dialects") + monkeypatch.setattr("sqlalchemy.dialects.registry", loader) + runpy.run_path(str(Path(__file__).parents[2] / "sqlalchemy" / "__init__.py")) + entry_points = { + entry_point.name: entry_point.load() + for entry_point in importlib.metadata.entry_points(group="sqlalchemy.dialects") + if entry_point.value.startswith("pyathena.") + } + assert entry_points + assert {name: load() for name, load in loader.impls.items()} == entry_points + @pytest.mark.parametrize("dialect_class", [AthenaDialect, AthenaAioDialect]) def test_type_compiler(self, dialect_class): # SQLAlchemy 2.0 builds the type compiler from type_compiler_cls. A legacy diff --git a/tests/sqlalchemy/__init__.py b/tests/sqlalchemy/__init__.py index 490f6a5fb..28d3c8bc6 100644 --- a/tests/sqlalchemy/__init__.py +++ b/tests/sqlalchemy/__init__.py @@ -7,10 +7,11 @@ from sqlalchemy.dialects import registry -registry.register("awsathena", "pyathena.sqlalchemy.base", "AthenaDialect") +registry.register("awsathena", "pyathena.sqlalchemy.rest", "AthenaRestDialect") registry.register("awsathena.rest", "pyathena.sqlalchemy.rest", "AthenaRestDialect") registry.register("awsathena.pandas", "pyathena.sqlalchemy.pandas", "AthenaPandasDialect") registry.register("awsathena.arrow", "pyathena.sqlalchemy.arrow", "AthenaArrowDialect") +registry.register("awsathena.polars", "pyathena.sqlalchemy.polars", "AthenaPolarsDialect") registry.register("awsathena.s3fs", "pyathena.sqlalchemy.s3fs", "AthenaS3FSDialect") registry.register("awsathena.aiorest", "pyathena.aio.sqlalchemy.rest", "AthenaAioRestDialect") registry.register("awsathena.aiopandas", "pyathena.aio.sqlalchemy.pandas", "AthenaAioPandasDialect")