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

Large diffs are not rendered by default.

Large diffs are not rendered by default.

Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,93 @@ def test_basic_setup(self, component_class, default_kwargs):
assert component.file_name == "test_output"
assert component.file_format == "csv"

def test_aws_fallback_inputs_are_optional(self, component_class):
"""Allow the canvas to defer AWS fallback-backed inputs to runtime."""
inputs = {component_input.name: component_input for component_input in component_class.inputs}

for input_name in ("aws_access_key_id", "aws_secret_access_key", "bucket_name"):
assert inputs[input_name].required is False

@pytest.mark.asyncio
async def test_save_to_aws_uses_environment_and_settings_fallbacks(self, component_class, monkeypatch):
"""Use resolved AWS values when the component credential inputs are empty."""
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "environment-access-key")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "environment-secret-key")
monkeypatch.setenv("AWS_SESSION_TOKEN", "environment-session-token")
monkeypatch.setenv("AWS_DEFAULT_REGION", "us-west-2")

settings_service = MagicMock()
settings_service.settings.object_storage_bucket_name = "settings-bucket"
mock_s3_client = MagicMock()

component = component_class()
component.set_attributes(
{
"input": Message(text="AWS fallback content"),
"file_name": "aws_report",
"storage_location": [{"name": "AWS"}],
"aws_format": "txt",
"aws_access_key_id": "",
"aws_secret_access_key": "",
"bucket_name": "",
"aws_region": "",
"s3_prefix": "reports",
}
)

with (
patch(
"lfx.components.files_and_knowledge.save_file.get_settings_service",
return_value=settings_service,
),
patch("boto3.client", return_value=mock_s3_client) as mock_boto3_client,
):
result = await component.save_to_file()

mock_boto3_client.assert_called_once_with(
"s3",
aws_access_key_id="environment-access-key",
aws_secret_access_key="environment-secret-key", # noqa: S106 # pragma: allowlist secret
aws_session_token="environment-session-token", # noqa: S106 # pragma: allowlist secret
region_name="us-west-2",
)
upload_args = mock_s3_client.upload_file.call_args.args
assert upload_args[1:] == ("settings-bucket", "reports/aws_report.txt")
assert result.text == "File successfully uploaded to s3://settings-bucket/reports/aws_report.txt"

@pytest.mark.asyncio
async def test_save_to_aws_does_not_mix_environment_token_with_component_credentials(
self, component_class, monkeypatch
):
"""Keep an environment session token out of component-supplied credentials."""
monkeypatch.setenv("AWS_SESSION_TOKEN", "environment-session-token")
mock_s3_client = MagicMock()

component = component_class()
component.set_attributes(
{
"input": Message(text="AWS component credential content"),
"file_name": "aws_report",
"storage_location": [{"name": "AWS"}],
"aws_format": "txt",
"aws_access_key_id": "component-access-key",
"aws_secret_access_key": "component-secret-key", # pragma: allowlist secret
"bucket_name": "component-bucket",
"aws_region": "us-west-2",
"s3_prefix": "reports",
}
)

with patch("boto3.client", return_value=mock_s3_client) as mock_boto3_client:
await component.save_to_file()

mock_boto3_client.assert_called_once_with(
"s3",
aws_access_key_id="component-access-key",
aws_secret_access_key="component-secret-key", # noqa: S106 # pragma: allowlist secret
region_name="us-west-2",
)

def test_get_input_type_dataframe(self, component_class):
"""Test input type detection for DataFrame."""
component = component_class()
Expand Down
18 changes: 9 additions & 9 deletions src/lfx/src/lfx/_assets/component_index.json

Large diffs are not rendered by default.

30 changes: 15 additions & 15 deletions src/lfx/src/lfx/components/files_and_knowledge/save_file.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,26 +144,26 @@ class SaveToFileComponent(Component):
SecretStrInput(
name="aws_access_key_id",
display_name="AWS Access Key ID",
info="AWS Access key ID.",
info="Optional. Falls back to the AWS_ACCESS_KEY_ID environment variable.",
show=_is_default_storage("AWS"),
advanced=not _is_default_storage("AWS"),
required=True,
required=False,
),
SecretStrInput(
name="aws_secret_access_key",
display_name="AWS Secret Key",
info="AWS Secret Key.",
info="Optional. Falls back to the AWS_SECRET_ACCESS_KEY environment variable.",
show=_is_default_storage("AWS"),
advanced=not _is_default_storage("AWS"),
required=True,
required=False,
),
StrInput(
name="bucket_name",
display_name="S3 Bucket Name",
info="Enter the name of the S3 bucket.",
info="Optional. Falls back to the configured object storage bucket.",
show=_is_default_storage("AWS"),
advanced=not _is_default_storage("AWS"),
required=True,
required=False,
),
StrInput(
name="aws_region",
Expand Down Expand Up @@ -687,19 +687,19 @@ async def _save_to_aws(self) -> Message:

import boto3

from lfx.base.data.cloud_storage_utils import create_s3_client, validate_aws_credentials

# Get AWS credentials from component inputs or fall back to environment variables
aws_access_key_id = getattr(self, "aws_access_key_id", None)
if aws_access_key_id and hasattr(aws_access_key_id, "get_secret_value"):
aws_access_key_id = aws_access_key_id.get_secret_value()
if not aws_access_key_id:
access_key_from_environment = not aws_access_key_id
if access_key_from_environment:
aws_access_key_id = os.getenv("AWS_ACCESS_KEY_ID")

aws_secret_access_key = getattr(self, "aws_secret_access_key", None)
if aws_secret_access_key and hasattr(aws_secret_access_key, "get_secret_value"):
aws_secret_access_key = aws_secret_access_key.get_secret_value()
if not aws_secret_access_key:
secret_key_from_environment = not aws_secret_access_key
if secret_key_from_environment:
aws_secret_access_key = os.getenv("AWS_SECRET_ACCESS_KEY")

bucket_name = getattr(self, "bucket_name", None)
Expand Down Expand Up @@ -728,15 +728,15 @@ async def _save_to_aws(self) -> Message:
)
raise ValueError(msg)

# Validate AWS credentials
validate_aws_credentials(self)

# Create S3 client
s3_client = create_s3_client(self)
# Create S3 client from the resolved component or fallback values
client_config: dict[str, Any] = {
"aws_access_key_id": str(aws_access_key_id),
"aws_secret_access_key": str(aws_secret_access_key),
}
if access_key_from_environment and secret_key_from_environment:
aws_session_token = os.getenv("AWS_SESSION_TOKEN")
if aws_session_token:
client_config["aws_session_token"] = aws_session_token

# Get region from component input, environment variable, or settings
aws_region = getattr(self, "aws_region", None)
Expand Down
Loading