diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 71b41c906..533639012 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -180,7 +180,7 @@ Column definitions in `CREATE TABLE` render `TIMESTAMP` for `DateTime` and for ` | 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 c58f88669..7dc4cced2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,7 +38,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 d956232c9..273681ebb 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -26,6 +26,7 @@ from pyathena.formatter import DefaultParameterFormatter from pyathena.sqlalchemy.base import AthenaDialect from pyathena.sqlalchemy.compiler import AthenaTypeCompiler +from pyathena.sqlalchemy.rest import AthenaRestDialect from pyathena.sqlalchemy.types import ( TINYINT, AthenaArray, @@ -110,6 +111,17 @@ def close(self): 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" + @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