"""Test Logger Class.""" from unittest.mock import patch from elasticsearch.exceptions import AuthorizationException import pytest from auditlogger.es_client_factory import create_es_client from auditlogger.exceptions import LoggingFailureException from auditlogger.logger import AuditLogger @patch("auditlogger.logger.DummyElasticSearch") def test_create_logger_no_key(mock_dummy_es): """Test creating a logger without credentials.""" audit_logger = AuditLogger([]) assert audit_logger assert mock_dummy_es.called @patch("builtins.print") def test_log_dummy_es(mocked_print, mock_record_event): """Test logging to the dummy ES prints the log.""" audit_logger = AuditLogger([]) audit_logger.log(mock_record_event) assert mocked_print.called mocked_print.assert_called_with( "Indexing to 'record-event':", mock_record_event.get_event_data() ) @patch("elasticsearch.transport.Transport.perform_request", return_value={}) def test_log_to_es(mock_request, mock_record_event, mock_aws_env, mock_hosts): """Test creating a logger backed by ES.""" AuditLogger(mock_hosts).log(mock_record_event) assert mock_request.call_count == 1 assert mock_request.call_args[0] == ("POST", "/record-event/_doc") @patch("elasticsearch.transport.Transport.perform_request", return_value={}) def test_log_to_es_and_os(mock_request, mock_record_event, mock_aws_env, mock_hosts, mock_os_host): """Test creating a logger backed by ES and OS.""" AuditLogger(mock_hosts, mock_os_host).log(mock_record_event) assert mock_request.call_count == 2 # for both ES and OS assert mock_request.call_args_list[0][0] == ("POST", "/record-event/_doc") # ES index name assert mock_request.call_args_list[1][0] == ("POST", "/record_event/_doc") # OS alias name def test_logger_invalid_host(mock_record_event, mock_aws_env, mock_hosts): """Test passing an invalid hostname.""" audit_logger = AuditLogger(mock_hosts) with pytest.raises(LoggingFailureException): audit_logger.log(mock_record_event) @patch("auditlogger.logger.create_es_client") def test_logger_multiple_attempts( mock_es_client_factory, mock_record_event, mock_hosts): """Test logging with retries.""" mock_es_client_factory.return_value.index.side_effect = Exception() audit_logger = AuditLogger(mock_hosts) with pytest.raises(LoggingFailureException): audit_logger.log(mock_record_event) assert mock_es_client_factory.return_value.index.call_count == 3 assert mock_es_client_factory.call_count == 1 @patch("auditlogger.logger.create_es_client") def test_logger_expired_session( mock_es_client_factory, mock_record_event, mock_hosts): """Test logging with expired boto session.""" mock_es_client_factory.return_value.index.side_effect = [ AuthorizationException(), None] audit_logger = AuditLogger(mock_hosts) audit_logger.log(mock_record_event) assert mock_es_client_factory.return_value.index.call_count == 2 assert mock_es_client_factory.call_count == 2 def test_logger_prepare_bulk(mock_hosts, mock_record_event): """Test the bulk data generator.""" logger = AuditLogger(mock_hosts) mock_data = [mock_record_event for i in range(0, 4)] index_key = 0 for item in logger.prepare_bulk(mock_data): assert item == { "_index": mock_data[index_key].ES_INDEX, "_type": "_doc", "_source": mock_data[index_key].get_event_data()} index_key += 1 def test_logger_prepare_bulk_os(mock_hosts, mock_os_host, mock_record_event): """Test the prepare_bulk_os generator.""" logger = AuditLogger(mock_hosts, mock_os_host) mock_data = [mock_record_event for i in range(0, 4)] index_key = 0 for item in logger.prepare_bulk_os(mock_data): assert item == { "_index": mock_data[index_key].OS_INDEX, "_source": mock_data[index_key].get_event_data()} index_key += 1 @patch("auditlogger.logger.bulk") def test_logger_log_bulk(mock_bulk, mock_hosts): """Test log_bulk.""" mock_data = ["foo", "bar", "baz", "qux"] logger = AuditLogger(mock_hosts) logger.log(mock_data) assert mock_bulk.call_count == 1 # only ES when no os_host is provided. @patch("auditlogger.logger.bulk") def test_logger_log_bulk_os(mock_bulk, mock_hosts, mock_os_host): """Test log_bulk.""" mock_data = ["foo", "bar", "baz", "qux"] logger = AuditLogger(mock_hosts, mock_os_host) logger.log(mock_data) assert mock_bulk.call_count == 2 # for both ES and OS @patch("auditlogger.es_client_factory.Elasticsearch") def test_create_es_client_with_user_pass(mock_es, mock_hosts_with_user_pass): """Test create_es_client.""" create_es_client(mock_hosts_with_user_pass) call_args = mock_es.call_args.kwargs["http_auth"] expected_args = ("bar", "baz") assert call_args == expected_args @patch("auditlogger.es_client_factory.Elasticsearch") @patch("auditlogger.es_client_factory.get_signature") def test_create_es_client_with_v4_signature(mock_get_signature, mock_es, mock_hosts): """Test create_es_client.""" create_es_client(mock_hosts) mock_get_signature.assert_called_once()