diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 7b75c0ea6..60fd19169 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -99,7 +99,7 @@ awsathena+aiorest://:@athena.{region_name}.amazonaws.com:443/{schema_name}?s3_st | Dialect | Driver | Schema | Cursor | |-----------|--------|------------------|------------------------| -| awsathena | | awsathena | DefaultCursor | +| awsathena | rest | awsathena | DefaultCursor | | awsathena | rest | awsathena+rest | DefaultCursor | | awsathena | pandas | awsathena+pandas | {ref}`pandas-cursor` | | awsathena | arrow | awsathena+arrow | {ref}`arrow-cursor` | diff --git a/pyproject.toml b/pyproject.toml index de1888064..1425c00c9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -36,7 +36,7 @@ documentation = "https://pyathena.dev/" issues = "https://github.com/pyathena-dev/PyAthena/issues" [project.entry-points."sqlalchemy.dialects"] -awsathena = "pyathena.sqlalchemy.base:AthenaDialect" +awsathena = "pyathena.sqlalchemy.rest:AthenaRestDialect" "awsathena.rest" = "pyathena.sqlalchemy.rest:AthenaRestDialect" "awsathena.pandas" = "pyathena.sqlalchemy.pandas:AthenaPandasDialect" "awsathena.arrow" = "pyathena.sqlalchemy.arrow:AthenaArrowDialect" diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 5dd7b1d52..2863fb77e 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -1,8 +1,11 @@ +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 urllib.parse import quote_plus import numpy as np @@ -15,7 +18,9 @@ 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.sqlalchemy.rest import AthenaRestDialect from pyathena.sqlalchemy.types import ( TINYINT, AthenaArray, @@ -50,6 +55,33 @@ def unique_s3tables_table_name(base: str) -> str: return f"{base}_{uuid.uuid4().hex[:8]}" +class TestAthenaDialect: + def test_bare_scheme_uses_rest_driver(self): + # The bare awsathena entry point resolves to the REST dialect, like + # awsathena+rest. Requires the package to be reinstalled (uv sync) so the + # installed entry point metadata matches pyproject.toml. + url = "awsathena://athena.us-west-2.amazonaws.com:443/default?s3_staging_dir=s3://bucket/path/" + bare = create_engine(url) + assert type(bare.dialect) is AthenaRestDialect + assert bare.driver == "rest" + 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 + + class TestSQLAlchemyAthena: @pytest.mark.parametrize( "engine", diff --git a/tests/sqlalchemy/__init__.py b/tests/sqlalchemy/__init__.py index 9a9dea779..b827aa6fe 100644 --- a/tests/sqlalchemy/__init__.py +++ b/tests/sqlalchemy/__init__.py @@ -1,9 +1,10 @@ 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")