From 797b62769576cf948cc279f6077273482dc45cf1 Mon Sep 17 00:00:00 2001 From: Hugo Pereira Brito <101209179+HugoPBrito@users.noreply.github.com> Date: Mon, 4 Nov 2024 14:55:47 +0100 Subject: [PATCH] feat(aws): add new service `transfer` (#5585) --- .../aws/services/transfer/__init__.py | 0 .../aws/services/transfer/transfer_client.py | 4 + .../aws/services/transfer/transfer_service.py | 69 ++++++++++++++++ .../transfer/transfer_service_test.py | 82 +++++++++++++++++++ 4 files changed, 155 insertions(+) create mode 100644 prowler/providers/aws/services/transfer/__init__.py create mode 100644 prowler/providers/aws/services/transfer/transfer_client.py create mode 100644 prowler/providers/aws/services/transfer/transfer_service.py create mode 100644 tests/providers/aws/services/transfer/transfer_service_test.py diff --git a/prowler/providers/aws/services/transfer/__init__.py b/prowler/providers/aws/services/transfer/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/prowler/providers/aws/services/transfer/transfer_client.py b/prowler/providers/aws/services/transfer/transfer_client.py new file mode 100644 index 0000000000..8f113f56f3 --- /dev/null +++ b/prowler/providers/aws/services/transfer/transfer_client.py @@ -0,0 +1,4 @@ +from prowler.providers.aws.services.transfer.transfer_service import Transfer +from prowler.providers.common.provider import Provider + +transfer_client = Transfer(Provider.get_global_provider()) diff --git a/prowler/providers/aws/services/transfer/transfer_service.py b/prowler/providers/aws/services/transfer/transfer_service.py new file mode 100644 index 0000000000..22bff6202a --- /dev/null +++ b/prowler/providers/aws/services/transfer/transfer_service.py @@ -0,0 +1,69 @@ +from enum import Enum +from typing import Dict, List + +from pydantic import BaseModel, Field + +from prowler.lib.logger import logger +from prowler.lib.scan_filters.scan_filters import is_resource_filtered +from prowler.providers.aws.lib.service.service import AWSService + + +class Transfer(AWSService): + def __init__(self, provider): + # Call AWSService's __init__ + super().__init__(__class__.__name__, provider) + self.servers = {} + self.__threading_call__(self._list_servers) + self.__threading_call__(self._describe_server, self.servers.values()) + + def _list_servers(self, regional_client): + logger.info("Transfer - Listing Transfer Servers...") + try: + list_servers_paginator = regional_client.get_paginator("list_servers") + for page in list_servers_paginator.paginate(): + for server in page["Servers"]: + arn = server["Arn"] + if not self.audit_resources or ( + is_resource_filtered(arn, self.audit_resources) + ): + self.servers[arn] = Server( + arn=arn, + id=server.get("ServerId", ""), + region=regional_client.region, + ) + except Exception as error: + logger.error( + f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" + ) + + def _describe_server(self, server): + logger.info(f"Transfer - Describing Server {server.id}...") + try: + server_description = ( + self.regional_clients[server.region] + .describe_server(ServerId=server.id) + .get("Server", {}) + ) + for protocol in server_description.get("Protocols", []): + server.protocols.append(Protocol(protocol)) + for tag in server_description.get("Tags", []): + server.tags.append(tag) + except Exception as error: + logger.error( + f"{server.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" + ) + + +class Protocol(Enum): + FTP = "FTP" + FTPS = "FTPS" + SFTP = "SFTP" + AS2 = "AS2" + + +class Server(BaseModel): + arn: str + id: str + region: str + protocols: List[Protocol] = Field(default_factory=list) + tags: List[Dict[str, str]] = Field(default_factory=list) diff --git a/tests/providers/aws/services/transfer/transfer_service_test.py b/tests/providers/aws/services/transfer/transfer_service_test.py new file mode 100644 index 0000000000..4c5b8db30b --- /dev/null +++ b/tests/providers/aws/services/transfer/transfer_service_test.py @@ -0,0 +1,82 @@ +from unittest.mock import patch + +import botocore +from moto import mock_aws + +from prowler.providers.aws.services.transfer.transfer_service import Protocol, Transfer +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_US_EAST_1, + set_mocked_aws_provider, +) + +make_api_call = botocore.client.BaseClient._make_api_call + +SERVER_ID = "SERVICE_MANAGED::s-01234567890abcdef" +SERVER_ARN = f"arn:aws:transfer:us-east-1:{AWS_ACCOUNT_NUMBER}:server/{SERVER_ID}" + + +def mock_make_api_call(self, operation_name, kwarg): + if operation_name == "ListServers": + return { + "Servers": [ + { + "Arn": f"arn:aws:transfer:{AWS_REGION_US_EAST_1}:{AWS_ACCOUNT_NUMBER}:server/{SERVER_ID}", + "ServerId": SERVER_ID, + } + ] + } + if operation_name == "DescribeServer": + return { + "Server": { + "Arn": SERVER_ARN, + "ServerId": SERVER_ID, + "Protocols": ["SFTP"], + "Tags": [{"key": "value"}], + } + } + return make_api_call(self, operation_name, kwarg) + + +class Test_transfer_service: + @mock_aws + def test_get_client(self): + aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) + transfer = Transfer(aws_provider) + assert ( + transfer.regional_clients[AWS_REGION_US_EAST_1].__class__.__name__ + == "Transfer" + ) + + @mock_aws + def test_get_session(self): + aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) + transfer = Transfer(aws_provider) + assert transfer.session.__class__.__name__ == "Session" + + @mock_aws + def test_get_service(self): + transfer = Transfer(set_mocked_aws_provider()) + assert transfer.service == "transfer" + + @patch("botocore.client.BaseClient._make_api_call", new=mock_make_api_call) + @mock_aws + def test_list_servers(self): + aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) + transfer = Transfer(aws_provider) + assert len(transfer.servers) == 1 + assert transfer.servers[SERVER_ARN].arn == SERVER_ARN + assert transfer.servers[SERVER_ARN].id == SERVER_ID + assert transfer.servers[SERVER_ARN].region == "us-east-1" + + @patch("botocore.client.BaseClient._make_api_call", new=mock_make_api_call) + @mock_aws + def test_describe_server(self): + aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) + transfer = Transfer(aws_provider) + assert transfer.servers[SERVER_ARN].arn == SERVER_ARN + assert transfer.servers[SERVER_ARN].id == SERVER_ID + assert len(transfer.servers[SERVER_ARN].protocols) == 1 + assert transfer.servers[SERVER_ARN].region == "us-east-1" + assert transfer.servers[SERVER_ARN].tags == [{"key": "value"}] + assert transfer.servers[SERVER_ARN].protocols[0] == Protocol.SFTP