diff --git a/docs/filesystem.md b/docs/filesystem.md index 8a46075db..dfc1effa2 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -41,7 +41,8 @@ fsspec.register_implementation("s3", s3fs.S3FileSystem, clobber=True) ## Basic usage -The filesystem can be constructed from a PyAthena connection, or directly with +The filesystem can be constructed from a PyAthena connection, whose S3 client it +then uses (see "S3 client" in [Usage](usage.md)), or directly with s3fs-compatible credential arguments: ```python diff --git a/docs/usage.md b/docs/usage.md index 7cddaff3f..492e7c11d 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -750,6 +750,34 @@ The connection builds one Glue client on first use from its session, region, and Pass `glue_metadata_fallback=False` to `connect()` to turn the fallback off. The Glue request does not carry the connection's workgroup; turn the fallback off where access depends on the workgroup, such as a workgroup enabled for IAM Identity Center. +## S3 client + +A connection builds one S3 client on first use and shares it with the result sets of its cursors, the `S3FileSystem` instances created with `connection=`, and the Spark cursors. +The result sets of `Cursor`, `DictCursor`, and their asynchronous versions do not use it, so these cursors build no S3 client. +`ArrowCursor`, and `PolarsCursor` for Parquet (`unload=True`) and chunked results, read the query results through their libraries' own S3 clients and use the shared client for their other S3 requests. + +The S3 client is built from the connection's session, region, and client arguments, with the botocore `config` merged with `s3_config`. +It does not use the connection's `endpoint_url` and `api_version`, which are Athena's, and neither do the S3 requests of `pyathena.pandas.util.to_sql()`. +To send S3 requests to another endpoint, use botocore's [service-specific endpoint settings](https://docs.aws.amazon.com/sdkref/latest/guide/feature-ss-endpoints.html), such as the `AWS_ENDPOINT_URL_S3` environment variable. +Options set in `s3_config` take precedence over those in `config` for the S3 client only: + +```python +from botocore.config import Config +from pyathena import connect + +conn = connect( + s3_staging_dir="s3://YOUR_S3_BUCKET/path/to/", + region_name="us-west-2", + s3_config=Config(max_pool_connections=50), +) +``` + +The client keeps up to `max_pool_connections` connections per host for reuse, 10 by default. +Concurrent requests beyond that open more connections, which urllib3 closes after use, logging a "Connection pool is full" warning. + +`Connection.close()` closes the network connections of the connection's Athena client, and of its Glue and S3 clients if they were built. +A client used after that opens new connections. + ## Environment variables Support [Boto3 environment variables](https://boto3.amazonaws.com/v1/documentation/api/latest/guide/configuration.html#using-environment-variables). diff --git a/pyathena/connection.py b/pyathena/connection.py index 826ff69d3..8b1435f9e 100644 --- a/pyathena/connection.py +++ b/pyathena/connection.py @@ -4,6 +4,7 @@ import logging import os +import threading import time from collections.abc import Callable from typing import ( @@ -132,6 +133,7 @@ def __init__( on_start_query_execution: Callable[[str], None] | None = ..., on_poll: OnPollCallback | None = ..., glue_metadata_fallback: bool = ..., + s3_config: Config | None = ..., **kwargs, ) -> None: ... @@ -165,6 +167,7 @@ def __init__( on_start_query_execution: Callable[[str], None] | None = ..., on_poll: OnPollCallback | None = ..., glue_metadata_fallback: bool = ..., + s3_config: Config | None = ..., **kwargs, ) -> None: ... @@ -197,6 +200,7 @@ def __init__( on_start_query_execution: Callable[[str], None] | None = None, on_poll: OnPollCallback | None = None, glue_metadata_fallback: bool = True, + s3_config: Config | None = None, **kwargs, ) -> None: """Initialize a new Athena database connection. @@ -243,6 +247,8 @@ def __init__( answer a throttled table-metadata, table-listing or database-listing request from the AWS Glue Data Catalog before retrying it. Defaults to True. + s3_config: Botocore Config options for the S3 client only, such as + ``max_pool_connections``. They are merged over ``config``. **kwargs: Additional arguments passed to boto3 Session and client. Raises: @@ -331,16 +337,16 @@ def __init__( **self._session_kwargs, ) - if not self.config.user_agent_extra or ( - pyathena.user_agent_extra not in self.config.user_agent_extra - ): - self.config.user_agent_extra = ( - f"{pyathena.user_agent_extra}" - f"{' ' + self.config.user_agent_extra if self.config.user_agent_extra else ''}" - ) + self._add_user_agent(self.config) + self.s3_config: Config = self.config.merge(s3_config) if s3_config else self.config + self._add_user_agent(self.s3_config) self._client = self._session.client( "athena", region_name=self.region_name, config=self.config, **self._client_kwargs ) + # Built on first use, once per connection even when several threads + # need it at the same time. + self._s3_client_lock = threading.Lock() + self._s3_client: BaseClient | None = None self._converter = converter self._formatter = formatter if formatter else DefaultParameterFormatter() self._retry_config = retry_config if retry_config else RetryConfig() @@ -356,6 +362,19 @@ def __init__( self._session, self.region_name, self.config, self._client_kwargs ) + @staticmethod + def _add_user_agent(config: Config) -> None: + """Add PyAthena's user agent to a botocore config unless it has it. + + Args: + config: The config to update in place. + """ + if not config.user_agent_extra or pyathena.user_agent_extra not in config.user_agent_extra: + config.user_agent_extra = ( + f"{pyathena.user_agent_extra}" + f"{' ' + config.user_agent_extra if config.user_agent_extra else ''}" + ) + def _assume_role( self, profile_name: str | None, @@ -477,6 +496,18 @@ def _client_kwargs(self) -> dict[str, Any]: """ return {k: v for k, v in self._kwargs.items() if k in self._CLIENT_PASSING_ARGS} + @property + def _s3_client_kwargs(self) -> dict[str, Any]: + """Get client keyword arguments for S3 client creation. + + Returns: + The client keyword arguments without Athena's ``endpoint_url`` + and ``api_version``. + """ + return { + k: v for k, v in self._client_kwargs.items() if k not in ("endpoint_url", "api_version") + } + @property def session(self) -> Session: """Get the boto3 session used for AWS API calls. @@ -495,6 +526,24 @@ def client(self) -> BaseClient: """ return self._client + @property + def s3_client(self) -> BaseClient: + """The S3 client shared by the connection's result sets, filesystems and Spark cursors. + + It is built on first use from the connection's session, region, + ``s3_config`` and client arguments, except Athena's ``endpoint_url`` + and ``api_version``. + """ + with self._s3_client_lock: + if self._s3_client is None: + self._s3_client = self._session.client( + "s3", + region_name=self.region_name, + config=self.s3_config, + **self._s3_client_kwargs, + ) + return self._s3_client + @property def retry_config(self) -> RetryConfig: """Get the retry configuration for AWS API calls. @@ -587,14 +636,19 @@ def cursor( def close(self) -> None: """Close the connection. - Closes the database connection. This method is provided for DB API 2.0 - compatibility. Since Athena connections are stateless, this method - currently does not perform any actual cleanup operations. + Closes the network connections of the Athena client and of the Glue and + S3 clients if they were built. A client used after this opens new + network connections. Note: This method is called automatically when using the connection as a context manager (with statement). """ + self._client.close() + self._glue.close() + with self._s3_client_lock: + if self._s3_client is not None: + self._s3_client.close() def commit(self) -> None: """Commit any pending transaction. diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index c78e1843a..acc7ae924 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -234,9 +234,10 @@ def __init__( """Create a filesystem for Amazon S3. Args: - connection: A PyAthena connection whose session, region, config and - retry policy the S3 client uses. Without one, the client is built - from s3fs-compatible arguments in ``kwargs``. + connection: A PyAthena connection whose S3 client + (``Connection.s3_client``) and retry policy the filesystem uses. + Without one, the client is built from s3fs-compatible arguments + in ``kwargs``. default_block_size: The block size for reads and writes; defaults to ``DEFAULT_BLOCK_SIZE``. default_cache_type: The fsspec cache type for reads; defaults to @@ -259,12 +260,7 @@ def __init__( """ super().__init__(*args, **kwargs) if connection: - client = connection.session.client( - "s3", - region_name=connection.region_name, - config=connection.config, - **connection._client_kwargs, - ) + client = connection.s3_client retry_config = connection.retry_config else: client = self._get_client_compatible_with_s3fs(**kwargs) diff --git a/pyathena/glue.py b/pyathena/glue.py index 7ac57e757..b429b20f7 100644 --- a/pyathena/glue.py +++ b/pyathena/glue.py @@ -93,6 +93,12 @@ def client(self) -> BaseClient: ) return self._client + def close(self) -> None: + """Close the network connections of the Glue client if it was built.""" + with self._lock: + if self._client is not None: + self._client.close() + @staticmethod def _catalog_request_kwargs(catalog_name: str | None) -> dict[str, str] | None: # AwsDataCatalog is the caller's default Glue catalog. An S3 Tables diff --git a/pyathena/pandas/util.py b/pyathena/pandas/util.py index 50689f79c..579461ccf 100644 --- a/pyathena/pandas/util.py +++ b/pyathena/pandas/util.py @@ -259,7 +259,7 @@ def to_sql( bucket_name, key_prefix = parse_output_location(location) bucket = conn.session.resource( - "s3", region_name=conn.region_name, **conn._client_kwargs + "s3", region_name=conn.region_name, **conn._s3_client_kwargs ).Bucket(bucket_name) cursor = conn.cursor() @@ -298,7 +298,7 @@ def to_sql( futures: list[concurrent.futures.Future[Any]] = [] session_kwargs = deepcopy(conn._session_kwargs) session_kwargs.update({"profile_name": conn.profile_name}) - client_kwargs = deepcopy(conn._client_kwargs) + client_kwargs = deepcopy(conn._s3_client_kwargs) client_kwargs.update({"region_name": conn.region_name}) partition_prefixes = [] if partitions: diff --git a/pyathena/result_set.py b/pyathena/result_set.py index 9ecc91876..2382b5b09 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -104,12 +104,6 @@ def __init__( self._hints_by_index[k] = v else: self._hints_by_name[k.lower()] = v - self._client = connection.session.client( - "s3", - region_name=connection.region_name, - config=connection.config, - **connection._client_kwargs, - ) self._metadata: tuple[dict[str, Any], ...] | None = None self._column_types: tuple[str, ...] | None = None @@ -773,7 +767,7 @@ def _get_content_length(self) -> int: bucket, key = parse_output_location(self.output_location) try: response = retry_api_call( - self._client.head_object, + self.connection.s3_client.head_object, config=self._retry_config, logger=_logger, Bucket=bucket, @@ -791,7 +785,7 @@ def _read_data_manifest(self) -> list[str]: bucket, key = parse_output_location(self.data_manifest_location) try: response = retry_api_call( - self._client.get_object, + self.connection.s3_client.get_object, config=self._retry_config, logger=_logger, Bucket=bucket, diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 0b9f2ce31..d6f12d77e 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -105,14 +105,9 @@ def __init__( self._calculation_id: str | None = None self._calculation_execution: AthenaCalculationExecution | None = None - # Created before the session so that a local failure cannot leave - # a newly started session behind. - self._client = self.connection.session.client( - "s3", - region_name=self.connection.region_name, - config=self.connection.config, - **self.connection._client_kwargs, - ) + # Taken before the session so that a failure to build the client + # cannot leave a newly started session behind. + self._client = self.connection.s3_client if session_id: if self._exists_session(session_id): diff --git a/pyathena/sqlalchemy/base.py b/pyathena/sqlalchemy/base.py index f408fdf06..dae07e9f7 100644 --- a/pyathena/sqlalchemy/base.py +++ b/pyathena/sqlalchemy/base.py @@ -306,6 +306,8 @@ def _create_connect_args(self, url: URL) -> dict[str, Any]: with contextlib.suppress(ValueError): verify = bool(strtobool(verify)) opts.update({"verify": verify}) + if "use_ssl" in opts: + opts.update({"use_ssl": bool(strtobool(opts["use_ssl"]))}) if "duration_seconds" in opts: opts.update({"duration_seconds": int(opts["duration_seconds"])}) if "poll_interval" in opts: diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 7ca3052a1..69a0b2b99 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -268,6 +268,30 @@ def test_result_set_file_system(self, pandas_cursor, chunksize): if not pandas_cursor.result_set.is_unload: assert pandas_cursor.result_set._csv_stream.closed + @pytest.mark.parametrize( + "pandas_cursor", + [{"endpoint_url": f"https://athena.{ENV.region_name}.amazonaws.com"}], + indirect=True, + ) + def test_athena_endpoint_url(self, pandas_cursor): + # GH-576: the S3 requests were sent to Athena's endpoint_url. + pandas_cursor.execute("SELECT * FROM one_row") + assert pandas_cursor.fetchall() == [(1,)] + + def test_result_sets_share_s3_client(self, pandas_cursor): + # GH-1011: each result set and its filesystem built their own S3 client, + # so every query opened new connections to S3. + conn = pandas_cursor.connection + session_client = conn.session.client + file_systems = [] + with patch.object(conn.session, "client", side_effect=session_client) as client: + for _ in range(2): + pandas_cursor.execute("SELECT * FROM one_row") + assert pandas_cursor.fetchall() == [(1,)] + file_systems.append(pandas_cursor.result_set._fs) + assert [c.args for c in client.call_args_list] == [("s3",)] + assert all(fs._client is conn.s3_client for fs in file_systems) + @pytest.mark.parametrize( ("query", "expected", "binary"), [ diff --git a/tests/pyathena/pandas/test_util.py b/tests/pyathena/pandas/test_util.py index c1ca82da7..d74ee4302 100644 --- a/tests/pyathena/pandas/test_util.py +++ b/tests/pyathena/pandas/test_util.py @@ -463,6 +463,21 @@ def test_to_sql_with_index(cursor): ] +@pytest.mark.parametrize( + "cursor", + [{"endpoint_url": f"https://athena.{ENV.region_name}.amazonaws.com"}], + indirect=True, +) +def test_to_sql_athena_endpoint_url(cursor): + # GH-576: the uploads were sent to Athena's endpoint_url. + df = pd.DataFrame({"col_int": np.int32([1])}) + table_name = f"""to_sql_{str(uuid.uuid4()).replace("-", "")}""" + location = f"{ENV.s3_staging_dir}{ENV.schema}/{table_name}/" + to_sql(df, table_name, cursor._connection, location, schema=ENV.schema, if_exists="fail") + cursor.execute(f"SELECT * FROM {table_name}") + assert cursor.fetchall() == [(1,)] + + def test_to_sql_with_partitions(cursor): df = pd.DataFrame( { diff --git a/tests/pyathena/spark/test_common.py b/tests/pyathena/spark/test_common.py index 5d16726e3..ad83598b3 100644 --- a/tests/pyathena/spark/test_common.py +++ b/tests/pyathena/spark/test_common.py @@ -9,7 +9,7 @@ import logging import threading import uuid -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, PropertyMock, patch import pytest from botocore.exceptions import ClientError @@ -480,7 +480,9 @@ def test_init_does_not_terminate_supplied_session(self, cursor_class): @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) def test_init_does_not_start_session_when_s3_client_fails(self, cursor_class): connection = _connection() - connection.session.client.side_effect = ValueError("Invalid S3 client configuration.") + type(connection).s3_client = PropertyMock( + side_effect=ValueError("Invalid S3 client configuration.") + ) with pytest.raises(ValueError, match=r"^Invalid S3 client configuration\.$"): _init_cursor(cursor_class, connection) diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 61d801e73..43ca5d511 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -16,6 +16,7 @@ import sqlalchemy from botocore.exceptions import ClientError from sqlalchemy import create_engine, engine_from_config, func, literal_column, select, text, types +from sqlalchemy.engine.url import make_url from sqlalchemy.exc import NoSuchTableError from sqlalchemy.sql import expression, type_coerce from sqlalchemy.sql.ddl import CreateTable @@ -140,6 +141,17 @@ def test_compliance_suite_registry_matches_entry_points(self, monkeypatch): assert entry_points assert {name: load() for name, load in loader.impls.items()} == entry_points + @pytest.mark.parametrize("dialect_class", [AthenaRestDialect, AthenaAioDialect]) + @pytest.mark.parametrize(("value", "expected"), [("false", False), ("true", True)]) + def test_conn_str_use_ssl(self, dialect_class, value, expected): + # The URL value is a string, and botocore treats "false" as true. + url = make_url( + "awsathena+rest://athena.us-west-2.amazonaws.com:443/default" + f"?s3_staging_dir=s3://bucket/path/&use_ssl={value}" + ) + _, opts = dialect_class().create_connect_args(url) + assert opts["use_ssl"] is expected + @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/pyathena/test_connection.py b/tests/pyathena/test_connection.py index 58f5a0608..7e6e42093 100644 --- a/tests/pyathena/test_connection.py +++ b/tests/pyathena/test_connection.py @@ -5,10 +5,16 @@ # # SPDX-License-Identifier: MIT +import os +import threading +from concurrent.futures import ThreadPoolExecutor from typing import Any +from unittest.mock import patch import pytest +from botocore.config import Config +import pyathena from pyathena.arrow.async_cursor import AsyncArrowCursor from pyathena.arrow.cursor import ArrowCursor from pyathena.async_cursor import AsyncCursor, AsyncDictCursor @@ -16,6 +22,7 @@ from pyathena.converter import DefaultTypeConverter from pyathena.cursor import Cursor, DictCursor from pyathena.error import ProgrammingError +from pyathena.filesystem.s3 import S3FileSystem from pyathena.pandas.async_cursor import AsyncPandasCursor from pyathena.pandas.cursor import PandasCursor from pyathena.polars.async_cursor import AsyncPolarsCursor @@ -55,6 +62,21 @@ def _connection(**kwargs: Any) -> Connection[Any]: ) +@pytest.fixture +def isolated_aws_config(monkeypatch, tmp_path): + """Hide the developer's AWS environment variables and config files. + + Args: + monkeypatch: The pytest monkeypatch fixture. + tmp_path: The pytest temporary directory, holding no config files. + """ + for key in list(os.environ): + if key.startswith("AWS_"): + monkeypatch.delenv(key) + monkeypatch.setenv("AWS_CONFIG_FILE", str(tmp_path / "config")) + monkeypatch.setenv("AWS_SHARED_CREDENTIALS_FILE", str(tmp_path / "credentials")) + + class TestConnection: @pytest.mark.parametrize( ("key", "configured", "explicit"), @@ -115,3 +137,93 @@ def test_cursor_arraysize_uncapped(self, cursor_class): def test_cursor_arraysize_not_positive(self, cursor_class): with pytest.raises(ProgrammingError): _connection().cursor(cursor_class, arraysize=0) + + def test_s3_client_built_once_across_threads(self): + conn = _connection() + created = [] + session_client = conn.session.client + + def contended_client(*args, **kwargs): + created.append((args, conn._s3_client_lock.locked())) + return session_client(*args, **kwargs) + + conn._session.client = contended_client + barrier = threading.Barrier(8) + + def get_client(_): + barrier.wait() + return conn.s3_client + + with ThreadPoolExecutor(max_workers=8) as executor: + clients = list(executor.map(get_client, range(8))) + + assert created == [(("s3",), True)] + assert all(client is clients[0] for client in clients) + assert clients[0].meta.service_model.service_name == "s3" + + def test_s3_client_leaves_out_athena_endpoint(self, isolated_aws_config): + # GH-576: Athena's endpoint_url (e.g. its VPC endpoint) was sent to S3. + conn = _connection( + endpoint_url="https://athena.us-east-1.amazonaws.com", + # Athena's API version, which S3 does not have. + api_version="2017-05-18", + ) + + assert conn.client.meta.endpoint_url == "https://athena.us-east-1.amazonaws.com" + assert conn.s3_client.meta.endpoint_url == "https://s3.amazonaws.com" + assert conn.s3_client.meta.service_model.api_version == "2006-03-01" + + def test_s3_client_uses_s3_endpoint_setting(self, isolated_aws_config, monkeypatch): + monkeypatch.setenv("AWS_ENDPOINT_URL_S3", "http://localhost:4566") + conn = _connection(endpoint_url="https://athena.us-east-1.amazonaws.com") + + assert conn.s3_client.meta.endpoint_url == "http://localhost:4566" + assert conn.client.meta.endpoint_url == "https://athena.us-east-1.amazonaws.com" + + def test_s3_filesystem_uses_connection_s3_client(self): + conn = _connection() + + fs = S3FileSystem(connection=conn, skip_instance_cache=True) + + assert fs._client is conn.s3_client + + def test_s3_config_defaults_to_config(self): + conn = _connection(config=Config(max_pool_connections=20)) + + assert conn.s3_config is conn.config + assert conn.s3_client.meta.config.max_pool_connections == 20 + + def test_s3_config_merged_over_config(self): + conn = _connection( + config=Config(connect_timeout=3, max_pool_connections=20), + s3_config=Config(max_pool_connections=50, user_agent_extra="s3-agent"), + ) + + s3_config = conn.s3_client.meta.config + assert s3_config.max_pool_connections == 50 + assert s3_config.connect_timeout == 3 + assert pyathena.user_agent_extra in s3_config.user_agent_extra + assert "s3-agent" in s3_config.user_agent_extra + assert conn.client.meta.config.max_pool_connections == 20 + assert "s3-agent" not in conn.client.meta.config.user_agent_extra + + def test_close_closes_built_clients(self): + conn = _connection() + with patch.object(conn.client, "close") as athena_close: + conn.close() + + athena_close.assert_called_once_with() + # Clients that were not used are not built to be closed. + assert conn._glue._client is None + assert conn._s3_client is None + + with ( + patch.object(conn.client, "close") as athena_close, + patch.object(conn._glue.client, "close") as glue_close, + patch.object(conn.s3_client, "close") as s3_close, + ): + conn.close() + + athena_close.assert_called_once_with() + glue_close.assert_called_once_with() + s3_close.assert_called_once_with() diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index 3f2485879..45000dfed 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -111,6 +111,14 @@ def start_query_execution(**kwargs): class TestCursor: + def test_builds_no_s3_client(self, cursor): + # GH-1011: the result set built an S3 client it never used. + session = cursor.connection.session + with patch.object(session, "client", side_effect=session.client) as client: + cursor.execute("SELECT * FROM one_row") + assert cursor.fetchall() == [(1,)] + client.assert_not_called() + def test_fetchone(self, cursor): cursor.execute("SELECT * FROM one_row") assert cursor.rowcount == -1