From 53dad647b631e002aa951b99bc5f442f9f15612e Mon Sep 17 00:00:00 2001 From: Benoit Gaudin Date: Tue, 14 Oct 2025 22:49:31 +0100 Subject: [PATCH] do not report aws_log_type tag if missing from destination config --- log_forwarder/data_retriever.py | 6 +++--- log_forwarder/destination_provider.py | 19 +++++++------------ log_forwarder/exceptions.py | 2 -- tests/test_destination_provider.py | 15 ++++++--------- 4 files changed, 16 insertions(+), 26 deletions(-) delete mode 100644 log_forwarder/exceptions.py diff --git a/log_forwarder/data_retriever.py b/log_forwarder/data_retriever.py index 46e74d0..415cbbf 100644 --- a/log_forwarder/data_retriever.py +++ b/log_forwarder/data_retriever.py @@ -16,7 +16,7 @@ class DataRetriever: def get_name(self): raise NotImplementedError() - def get_collection_type(self): + def get_source_type(self): raise NotImplementedError() def get_data(self): @@ -37,7 +37,7 @@ def __init__(self, config: Config, bucket_name, s3_key): self.s3_client = S3Client(config.filepath) logger.debug('type=%s, bucket=%s, s3_key=%s', self.__class__, self.src_bucket_name, self.src_key) - def get_collection_type(self): + def get_source_type(self): return 's3' def get_name(self): @@ -122,7 +122,7 @@ class CloudwatchDataRetriever(DataRetriever): def get_name(self): return 'Cloudwatch' - def get_collection_type(self): + def get_source_type(self): return 'cloudwatch' def __init__(self, config: Config): diff --git a/log_forwarder/destination_provider.py b/log_forwarder/destination_provider.py index 41b5f5e..5f4adbd 100644 --- a/log_forwarder/destination_provider.py +++ b/log_forwarder/destination_provider.py @@ -1,6 +1,5 @@ from data_retriever import DataRetriever from config import DestinationConfig, CLOUDWATCH_LOG_TYPE -from exceptions import LogTypeMissingException class DestinationProvider: @@ -9,20 +8,16 @@ def __init__(self, dest_config: DestinationConfig, data_retriever: DataRetriever self._config: DestinationConfig = dest_config self._data_retriever: DataRetriever = data_retriever - def get_type(self, data_id): - if (self._config.destination_config is not None and data_id in self._config.get_keys() and - self._config.get_log_type(data_id) is not None): - return self._config.get_log_type(data_id) - if self._data_retriever.get_collection_type() == 'cloudwatch': + def get_log_type(self, data_id): + if self._data_retriever.get_source_type() == 'cloudwatch': return CLOUDWATCH_LOG_TYPE - raise LogTypeMissingException('Log type required in configuration. data_type=%s, data_id=%s', - self._data_retriever.get_collection_type(), data_id) + return self._config.get_log_type(data_id) def get_dataset(self, data_id): if (self._config.destination_config is not None and data_id in self._config.get_keys() and self._config.get_dataset(data_id) is not None): return self._config.get_dataset(data_id) - if self._data_retriever.get_collection_type() == 'cloudwatch': + if self._data_retriever.get_source_type() == 'cloudwatch': return data_id return None @@ -30,12 +25,12 @@ def get_collection(self, data_id): if (self._config.destination_config is not None and data_id in self._config.get_keys() and self._config.get_collection(data_id) is not None): return self._config.get_collection(data_id) - if self._data_retriever.get_collection_type() == 'cloudwatch': + if self._data_retriever.get_source_type() == 'cloudwatch': return self._config.get_cloudwatch_default_collection() return None def get_dataset_tags(self, data_id): - log_type = self.get_type(data_id) - tags = {'aws_log_type': log_type if log_type is not None else 'unknown'} + log_type = self.get_log_type(data_id) + tags = {'aws_log_type': log_type} if log_type is not None else {} tags.update(self._config.get_dataset_tags(data_id)) return tags diff --git a/log_forwarder/exceptions.py b/log_forwarder/exceptions.py deleted file mode 100644 index 1db666b..0000000 --- a/log_forwarder/exceptions.py +++ /dev/null @@ -1,2 +0,0 @@ -class LogTypeMissingException(Exception): - pass diff --git a/tests/test_destination_provider.py b/tests/test_destination_provider.py index 15f5a6c..ec1f7d0 100644 --- a/tests/test_destination_provider.py +++ b/tests/test_destination_provider.py @@ -1,11 +1,10 @@ import json import base64 import tempfile -from exceptions import LogTypeMissingException from config import DestinationConfig, Config, CLOUDWATCH_LOG_TYPE -from data_retriever import CloudwatchDataRetriever, DataRetriever, S3DataRetriever, CustomS3Retriever, \ - DataRetrieverFactory, logger +from data_retriever import (CloudwatchDataRetriever, DataRetriever, S3DataRetriever, CustomS3Retriever, + DataRetrieverFactory) from destination_provider import DestinationProvider import pytest @@ -14,7 +13,7 @@ def _get_data(data_retriever, log_group_name): data_retriever.log_group_name = log_group_name -def test_s3_requires_log_type_in_config(): +def test_s3_no_log_type_in_config(): bucket_name = 'my_bucket' s3_key = 'my_key' with tempfile.NamedTemporaryFile() as f: @@ -23,8 +22,7 @@ def test_s3_requires_log_type_in_config(): dest_config = DestinationConfig() data_retriever = S3DataRetriever(config, bucket_name, s3_key) destination_provider = DestinationProvider(dest_config, data_retriever) - with pytest.raises(LogTypeMissingException) as _: - destination_provider.get_type(data_id) + assert destination_provider.get_log_type(data_id) is None def test_s3_custom_path(): @@ -49,10 +47,10 @@ def test_cloudwatch_no_config(monkeypatch): data_retriever = CloudwatchDataRetriever(config) monkeypatch.setattr(DataRetriever, 'get_data', lambda: _get_data(data_retriever, log_group_name)) destination_provider = DestinationProvider(dest_config, data_retriever) - assert destination_provider.get_type(log_group_name) == CLOUDWATCH_LOG_TYPE assert destination_provider.get_dataset(log_group_name) == log_group_name assert destination_provider.get_collection(log_group_name) is None - assert destination_provider.get_dataset_tags(log_group_name) == {'aws_log_type': 'cloudwatch_log'} + assert destination_provider.get_log_type(log_group_name) == CLOUDWATCH_LOG_TYPE + assert destination_provider.get_dataset_tags(log_group_name) == {'aws_log_type': CLOUDWATCH_LOG_TYPE} def test_cloudwatch_default_collection(monkeypatch): @@ -80,7 +78,6 @@ def test_cloudwatch_config_takes_precedence(monkeypatch): # mock _get_data in order to associate log group name to the data retriever monkeypatch.setattr(DataRetriever, 'get_data', lambda: _get_data(data_retriever, log_group_name_from_config)) destination_provider = DestinationProvider(dest_config, data_retriever) - assert destination_provider.get_type(log_group_name) == CLOUDWATCH_LOG_TYPE assert destination_provider.get_dataset(log_group_name) == log_group_name_from_config assert destination_provider.get_collection(log_group_name) == log_set_from_config