diff --git a/docs/pandas.md b/docs/pandas.md index bdb0f749..68d050ef 100644 --- a/docs/pandas.md +++ b/docs/pandas.md @@ -112,6 +112,8 @@ print(cursor.fetchall()) `to_sql` writes the data as Parquet with pyarrow, so it requires `pip install PyAthena[Pandas,Arrow]`. Conversion to Parquet and upload to S3 use [ThreadPoolExecutor](https://docs.python.org/3/library/concurrent.futures.html#threadpoolexecutor) by default. It is also possible to use [ProcessPoolExecutor](https://docs.python.org/3/library/concurrent.futures.html#processpoolexecutor). +The S3 requests use the connection's `s3_config` (see "S3 client" in [Usage](usage.md)) and credentials. +The upload workers resolve the credentials themselves, except for a connection given a `session` and no explicit keys, whose session's credentials they receive once, when the uploads start. ```python import pandas as pd diff --git a/pyathena/connection.py b/pyathena/connection.py index 8b1435f9..6d52074b 100644 --- a/pyathena/connection.py +++ b/pyathena/connection.py @@ -295,6 +295,9 @@ def __init__( if self.s3_staging_dir and not self.s3_staging_dir.endswith("/"): self.s3_staging_dir = f"{self.s3_staging_dir}/" + # Whether the session was given rather than built from the arguments, + # which then do not reproduce its credentials. + self._session_given = bool(session) if session: self._session = session else: diff --git a/pyathena/pandas/util.py b/pyathena/pandas/util.py index 579461cc..899a2344 100644 --- a/pyathena/pandas/util.py +++ b/pyathena/pandas/util.py @@ -215,6 +215,11 @@ def to_sql( as Parquet files to S3 and executing the appropriate DDL statements. Supports partitioning, compression, and parallel uploads. + The S3 requests use the connection's ``s3_config`` and credentials. The + upload workers resolve the credentials themselves, except for a connection + given a ``session`` and no explicit keys, whose session's credentials they + receive once, when the uploads start. + Args: df: The DataFrame to write to Athena. name: Name of the table to create. @@ -259,7 +264,7 @@ def to_sql( bucket_name, key_prefix = parse_output_location(location) bucket = conn.session.resource( - "s3", region_name=conn.region_name, **conn._s3_client_kwargs + "s3", region_name=conn.region_name, config=conn.s3_config, **conn._s3_client_kwargs ).Bucket(bucket_name) cursor = conn.cursor() @@ -296,10 +301,29 @@ def to_sql( reset_index(df, index_label) with executor_class(max_workers=max_workers) as e: futures: list[concurrent.futures.Future[Any]] = [] + # The workers build their own sessions from these arguments, which + # resolve the connection's credentials again unless the connection was + # given its session. Then they get its credentials as of now, unless + # explicit keys or a botocore session, which boto3 would set them on, + # take precedence. session_kwargs = deepcopy(conn._session_kwargs) session_kwargs.update({"profile_name": conn.profile_name}) + if ( + conn._session_given + and not conn._s3_client_kwargs.get("aws_access_key_id") + and not session_kwargs.get("botocore_session") + and (credentials := conn.session.get_credentials()) + ): + frozen_credentials = credentials.get_frozen_credentials() + session_kwargs.update( + { + "aws_access_key_id": frozen_credentials.access_key, + "aws_secret_access_key": frozen_credentials.secret_key, + "aws_session_token": frozen_credentials.token, + } + ) client_kwargs = deepcopy(conn._s3_client_kwargs) - client_kwargs.update({"region_name": conn.region_name}) + client_kwargs.update({"region_name": conn.region_name, "config": conn.s3_config}) partition_prefixes = [] if partitions: for keys, group in df.groupby(by=partitions, observed=True): diff --git a/tests/pyathena/pandas/test_util.py b/tests/pyathena/pandas/test_util.py index d74ee430..4564f3fc 100644 --- a/tests/pyathena/pandas/test_util.py +++ b/tests/pyathena/pandas/test_util.py @@ -1,13 +1,20 @@ +import contextlib import textwrap import uuid +from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor from datetime import date, datetime from decimal import Decimal +from multiprocessing import get_context +from unittest.mock import patch import numpy as np import pandas as pd import pytest +from boto3.session import Session +from botocore.config import Config from pyathena import OperationalError +from pyathena.pandas import util from pyathena.pandas.util import ( as_pandas, generate_ddl, @@ -16,6 +23,7 @@ to_sql, ) from tests import ENV +from tests.pyathena.conftest import connect def test_get_chunks(): @@ -478,6 +486,83 @@ def test_to_sql_athena_endpoint_url(cursor): assert cursor.fetchall() == [(1,)] +class SpawnProcessPoolExecutor(ProcessPoolExecutor): + """Process pool whose workers exit with it and keep no environment for later pools.""" + + def __init__(self, max_workers=None): + super().__init__(max_workers, mp_context=get_context("spawn")) + + +@pytest.mark.parametrize( + ("executor_class", "connect_kwargs"), + [ + (ThreadPoolExecutor, {}), + (SpawnProcessPoolExecutor, {}), + # As SQLAlchemy passes them without credentials in the URL. + (ThreadPoolExecutor, {"aws_access_key_id": None, "aws_secret_access_key": None}), + ], +) +def test_to_sql_session_credentials_and_s3_config(monkeypatch, executor_class, connect_kwargs): + # GH-1067: the upload workers used the default credential chain instead of + # the credentials of connect(session=...), and no S3 request used s3_config. + credentials = Session().get_credentials().get_frozen_credentials() + session = Session( + aws_access_key_id=credentials.access_key, + aws_secret_access_key=credentials.secret_key, + aws_session_token=credentials.token, + region_name=ENV.region_name, + ) + # The default credential chain, as the workers would resolve it, fails. + for key in ("AWS_PROFILE", "AWS_DEFAULT_PROFILE", "AWS_SESSION_TOKEN"): + monkeypatch.delenv(key, raising=False) + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "AKIAIOSFODNN7EXAMPLE") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "invalid") + df = pd.DataFrame({"col_int": np.int32([1, 2])}) + table_name = f"""to_sql_{str(uuid.uuid4()).replace("-", "")}""" + location = f"{ENV.s3_staging_dir}{ENV.schema}/{table_name}/" + resource = Session.resource + with ( + contextlib.closing( + connect( + schema_name=ENV.schema, + session=session, + s3_config=Config(max_pool_connections=37), + **connect_kwargs, + ) + ) as conn, + patch.object(Session, "resource", autospec=True, side_effect=resource) as resources, + ): + to_sql( + df, + table_name, + conn, + location, + schema=ENV.schema, + chunksize=1, + executor_class=executor_class, + max_workers=2, + ) + cursor = conn.cursor() + cursor.execute(f"SELECT * FROM {table_name} ORDER BY col_int") + assert cursor.fetchall() == [(1,), (2,)] + # The bucket resource, and with threads also the workers' resources. + expected = 1 if executor_class is SpawnProcessPoolExecutor else 3 + assert len(resources.call_args_list) == expected + assert all(c.kwargs["config"].max_pool_connections == 37 for c in resources.call_args_list) + + +def test_to_sql_workers_resolve_credentials(cursor): + # Without a given session, the workers resolve the credentials themselves, + # so refreshable credentials stay refreshable. + 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}/" + with patch.object(util, "to_parquet", side_effect=util.to_parquet) as to_parquet: + to_sql(df, table_name, cursor._connection, location, schema=ENV.schema) + session_kwargs = to_parquet.call_args.args[4] + assert "aws_secret_access_key" not in session_kwargs + + def test_to_sql_with_partitions(cursor): df = pd.DataFrame( {