Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
Prev Previous commit
Next Next commit
fix(spark): use AWS_ENDPOINT_URL, support path-style addressing, fix …
…linting

Signed-off-by: abhijeet-dhumal <abhijeetdhumal652@gmail.com>
  • Loading branch information
abhijeet-dhumal authored and ntkathole committed Apr 27, 2026
commit c7b74dbf4050c0540f21daa11e7676ec359167b0
27 changes: 21 additions & 6 deletions sdk/python/feast/infra/compute_engines/spark/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,7 @@ def _ensure_s3a_event_log_dir(spark_config: Dict[str, str]) -> None:

endpoint = spark_config.get(
"spark.hadoop.fs.s3a.endpoint",
os.environ.get("FEAST_S3A_ENDPOINT", ""),
os.environ.get("AWS_ENDPOINT_URL", ""),
)
access_key = spark_config.get(
"spark.hadoop.fs.s3a.access.key",
Expand All @@ -58,22 +58,37 @@ def _ensure_s3a_event_log_dir(spark_config: Dict[str, str]) -> None:
"spark.hadoop.fs.s3a.secret.key",
os.environ.get("AWS_SECRET_ACCESS_KEY", ""),
)
session_token = spark_config.get(
"spark.hadoop.fs.s3a.session.token",
os.environ.get("AWS_SESSION_TOKEN", ""),
) or None
session_token = (
spark_config.get(
"spark.hadoop.fs.s3a.session.token",
os.environ.get("AWS_SESSION_TOKEN", ""),
)
or None
)

try:
if boto3 is None:
raise ImportError("boto3 is not installed")

addressing_style = (
"path"
if spark_config.get(
"spark.hadoop.fs.s3a.path.style.access", "false"
).lower()
== "true"
else "auto"
)

s3 = boto3.client(
"s3",
endpoint_url=endpoint if endpoint else None,
aws_access_key_id=access_key or None,
aws_secret_access_key=secret_key or None,
aws_session_token=session_token,
config=BotoConfig(signature_version="s3v4"),
config=BotoConfig(
signature_version="s3v4",
s3={"addressing_style": addressing_style},
),
)
resp = s3.list_objects_v2(Bucket=bucket, Prefix=prefix, MaxKeys=1)
if resp.get("KeyCount", 0) == 0:
Expand Down
2 changes: 1 addition & 1 deletion sdk/python/tests/component/spark/test_compute.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
from datetime import timedelta
from typing import cast
from unittest.mock import MagicMock, patch
from unittest.mock import MagicMock

import pytest
from pyspark.sql import DataFrame
Expand Down
98 changes: 86 additions & 12 deletions sdk/python/tests/component/spark/test_spark_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -91,12 +91,8 @@ def test_ensure_s3a_event_log_dir_bucket_root_no_trailing_slash(mock_boto3):

_ensure_s3a_event_log_dir(_base_conf("s3a://my-bucket"))

s3.list_objects_v2.assert_called_once_with(
Bucket="my-bucket", Prefix="", MaxKeys=1
)
s3.put_object.assert_called_once_with(
Bucket="my-bucket", Key=".keep", Body=b""
)
s3.list_objects_v2.assert_called_once_with(Bucket="my-bucket", Prefix="", MaxKeys=1)
s3.put_object.assert_called_once_with(Bucket="my-bucket", Key=".keep", Body=b"")


@patch(BOTOCONFIG_PATH, MagicMock())
Expand All @@ -109,12 +105,8 @@ def test_ensure_s3a_event_log_dir_bucket_root_trailing_slash(mock_boto3):

_ensure_s3a_event_log_dir(_base_conf("s3a://my-bucket/"))

s3.list_objects_v2.assert_called_once_with(
Bucket="my-bucket", Prefix="", MaxKeys=1
)
s3.put_object.assert_called_once_with(
Bucket="my-bucket", Key=".keep", Body=b""
)
s3.list_objects_v2.assert_called_once_with(Bucket="my-bucket", Prefix="", MaxKeys=1)
s3.put_object.assert_called_once_with(Bucket="my-bucket", Key=".keep", Body=b"")


# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -198,3 +190,85 @@ def test_ensure_s3a_event_log_dir_no_credentials_passes_none(mock_boto3):
assert kwargs.kwargs["aws_access_key_id"] is None
assert kwargs.kwargs["aws_secret_access_key"] is None
assert kwargs.kwargs["aws_session_token"] is None


# ---------------------------------------------------------------------------
# Path-style addressing (MinIO / S3-compatible)
# ---------------------------------------------------------------------------


@patch(BOTOCONFIG_PATH)
@patch(BOTO3_PATH)
def test_ensure_s3a_event_log_dir_path_style_when_enabled(mock_boto3, mock_config_cls):
"""spark.hadoop.fs.s3a.path.style.access=true -> addressing_style='path'."""
s3 = MagicMock()
mock_boto3.client.return_value = s3
s3.list_objects_v2.return_value = {"KeyCount": 1}

conf = {
**_base_conf("s3a://my-bucket/logs/"),
"spark.hadoop.fs.s3a.path.style.access": "true",
}
_ensure_s3a_event_log_dir(conf)

mock_config_cls.assert_called_once()
config_kwargs = mock_config_cls.call_args
assert config_kwargs.kwargs["s3"] == {"addressing_style": "path"}


@patch(BOTOCONFIG_PATH)
@patch(BOTO3_PATH)
def test_ensure_s3a_event_log_dir_virtual_hosted_style_by_default(
mock_boto3, mock_config_cls
):
"""No path.style.access config -> addressing_style='auto'."""
s3 = MagicMock()
mock_boto3.client.return_value = s3
s3.list_objects_v2.return_value = {"KeyCount": 1}

_ensure_s3a_event_log_dir(_base_conf("s3a://my-bucket/logs/"))

mock_config_cls.assert_called_once()
config_kwargs = mock_config_cls.call_args
assert config_kwargs.kwargs["s3"] == {"addressing_style": "auto"}


# ---------------------------------------------------------------------------
# Endpoint env var fallback (AWS_ENDPOINT_URL)
# ---------------------------------------------------------------------------


@patch.dict("os.environ", {"AWS_ENDPOINT_URL": "http://localhost:9000"}, clear=True)
@patch(BOTOCONFIG_PATH, MagicMock())
@patch(BOTO3_PATH)
def test_ensure_s3a_event_log_dir_endpoint_from_env(mock_boto3):
"""AWS_ENDPOINT_URL env var is used when spark config has no endpoint."""
s3 = MagicMock()
mock_boto3.client.return_value = s3
s3.list_objects_v2.return_value = {"KeyCount": 1}

conf = {
"spark.eventLog.enabled": "true",
"spark.eventLog.dir": "s3a://my-bucket/logs/",
}
_ensure_s3a_event_log_dir(conf)

mock_boto3.client.assert_called_once()
kwargs = mock_boto3.client.call_args
assert kwargs.kwargs["endpoint_url"] == "http://localhost:9000"


@patch.dict("os.environ", {"AWS_ENDPOINT_URL": "http://env-endpoint:9000"}, clear=True)
@patch(BOTOCONFIG_PATH, MagicMock())
@patch(BOTO3_PATH)
def test_ensure_s3a_event_log_dir_spark_endpoint_over_env(mock_boto3):
"""spark.hadoop.fs.s3a.endpoint takes precedence over AWS_ENDPOINT_URL."""
s3 = MagicMock()
mock_boto3.client.return_value = s3
s3.list_objects_v2.return_value = {"KeyCount": 1}

_ensure_s3a_event_log_dir(_base_conf("s3a://my-bucket/logs/"))

mock_boto3.client.assert_called_once()
kwargs = mock_boto3.client.call_args
assert kwargs.kwargs["endpoint_url"] == "http://minio:9000"