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
2 changes: 1 addition & 1 deletion docs/sqlalchemy.md
Original file line number Diff line number Diff line change
Expand Up @@ -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` |
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
12 changes: 12 additions & 0 deletions tests/pyathena/sqlalchemy/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
Loading