"""Unit tests for general logging utilities module.""" import logging import sys from io import StringIO from unittest.mock import patch import pytest from abacus_common_logic.utils.logging import ( LOGGER_LEVEL, TaskContextFilter, TaskContextFormatter, create_handler, task_id_var, ) @pytest.fixture(autouse=True) def reset_task_id(): """Reset task_id_var before and after each test.""" # Reset before test token = task_id_var.set('') yield # Reset after test task_id_var.reset(token) class TestTaskContextFilter: """Tests for TaskContextFilter.""" def test_filter_injects_task_id_default(self): """Test that filter injects default task_id (empty string).""" filter_ = TaskContextFilter() record = logging.LogRecord( name='test', level=logging.INFO, pathname=__file__, lineno=10, msg='test message', args=(), exc_info=None, ) # Default task_id should be empty string result = filter_.filter(record) assert result is True assert hasattr(record, 'task_id') assert record.task_id == '' def test_filter_injects_custom_task_id(self): """Test that filter injects custom task_id from context.""" filter_ = TaskContextFilter() record = logging.LogRecord( name='test', level=logging.INFO, pathname=__file__, lineno=10, msg='test message', args=(), exc_info=None, ) task_id_var.set('task-123') result = filter_.filter(record) assert result is True assert record.task_id == 'task-123' def test_filter_updates_on_context_change(self): """Test that filter reflects changes to task_id context.""" filter_ = TaskContextFilter() # First record with task-abc task_id_var.set('task-abc') record1 = logging.LogRecord( name='test', level=logging.INFO, pathname=__file__, lineno=10, msg='message 1', args=(), exc_info=None, ) filter_.filter(record1) assert record1.task_id == 'task-abc' # Second record with task-xyz task_id_var.set('task-xyz') record2 = logging.LogRecord( name='test', level=logging.INFO, pathname=__file__, lineno=11, msg='message 2', args=(), exc_info=None, ) filter_.filter(record2) assert record2.task_id == 'task-xyz' def test_filter_always_returns_true(self): """Test that filter always returns True to process all records.""" filter_ = TaskContextFilter() record = logging.LogRecord( name='test', level=logging.DEBUG, pathname=__file__, lineno=10, msg='test', args=(), exc_info=None, ) # Should always return True regardless of record content assert filter_.filter(record) is True class TestTaskContextFormatter: """Tests for TaskContextFormatter.""" def test_format_without_task_id(self): """Test formatting when task_id is empty.""" formatter = TaskContextFormatter() record = logging.LogRecord( name='test_logger', level=logging.INFO, pathname=__file__, lineno=10, msg='test message', args=(), exc_info=None, ) record.task_id = '' # Empty task_id formatted = formatter.format(record) # Should not include task_id brackets when empty assert '[test_logger]' in formatted assert '[INFO]' in formatted assert 'test message' in formatted assert '[]' not in formatted # No empty brackets def test_format_with_task_id(self): """Test formatting when task_id is present.""" formatter = TaskContextFormatter() record = logging.LogRecord( name='test_logger', level=logging.INFO, pathname=__file__, lineno=10, msg='test message', args=(), exc_info=None, ) record.task_id = 'task-456' formatted = formatter.format(record) # Should include task_id in brackets assert '[test_logger]' in formatted assert '[INFO]' in formatted assert '[task-456]' in formatted assert 'test message' in formatted def test_format_with_different_log_levels(self): """Test formatting with different log levels.""" formatter = TaskContextFormatter() for level, level_name in [ (logging.DEBUG, 'DEBUG'), (logging.INFO, 'INFO'), (logging.WARNING, 'WARNING'), (logging.ERROR, 'ERROR'), (logging.CRITICAL, 'CRITICAL'), ]: record = logging.LogRecord( name='test', level=level, pathname=__file__, lineno=10, msg='message', args=(), exc_info=None, ) record.task_id = '' formatted = formatter.format(record) assert f'[{level_name}]' in formatted def test_format_includes_timestamp(self): """Test that formatting includes timestamp.""" formatter = TaskContextFormatter() record = logging.LogRecord( name='test', level=logging.INFO, pathname=__file__, lineno=10, msg='message', args=(), exc_info=None, ) record.task_id = '' formatted = formatter.format(record) # Should contain timestamp pattern (rough check) # Format is [name][timestamp][level] message # Timestamp should be between name and level assert formatted.count('[') >= 3 # At least [name], [timestamp], [level] def test_format_message_with_args(self): """Test formatting with message arguments.""" formatter = TaskContextFormatter() record = logging.LogRecord( name='test', level=logging.INFO, pathname=__file__, lineno=10, msg='User %s logged in with ID %d', args=('john', 123), exc_info=None, ) record.task_id = 'login-task' formatted = formatter.format(record) assert 'User john logged in with ID 123' in formatted assert '[login-task]' in formatted class TestCreateHandler: """Tests for create_handler factory function.""" def test_create_handler_default_level(self): """Test create_handler with default level.""" handler = create_handler() assert isinstance(handler, logging.StreamHandler) assert handler.level == LOGGER_LEVEL assert isinstance(handler.formatter, TaskContextFormatter) assert handler.stream is sys.stdout def test_create_handler_custom_level(self): """Test create_handler with custom level.""" handler = create_handler(level=logging.DEBUG) assert handler.level == logging.DEBUG def test_create_handler_has_filter(self): """Test that created handler has TaskContextFilter.""" handler = create_handler() # Check that handler has at least one filter filters = [f for f in handler.filters if isinstance(f, TaskContextFilter)] assert len(filters) == 1 def test_create_handler_creates_independent_instances(self): """Test that each call creates a new handler instance.""" handler1 = create_handler() handler2 = create_handler() assert handler1 is not handler2 assert handler1.formatter is not handler2.formatter def test_create_handler_integration(self): """Test full integration of created handler with logger.""" # Create a logger with the handler test_logger = logging.getLogger('test_integration') test_logger.handlers.clear() test_logger.setLevel(logging.INFO) test_logger.propagate = False # Capture output stream = StringIO() handler = logging.StreamHandler(stream) handler.setLevel(logging.INFO) handler.setFormatter(TaskContextFormatter()) handler.addFilter(TaskContextFilter()) test_logger.addHandler(handler) # Test without task_id task_id_var.set('') test_logger.info('Message without task') output_no_task = stream.getvalue() assert 'Message without task' in output_no_task assert '[]' not in output_no_task # No empty brackets # Test with task_id stream.truncate(0) stream.seek(0) task_id_var.set('test-task-id') test_logger.info('Message with task') output_with_task = stream.getvalue() assert 'Message with task' in output_with_task assert '[test-task-id]' in output_with_task class TestTaskIdVar: """Tests for task_id_var context variable.""" def test_task_id_var_default(self): """Test task_id_var has empty string default.""" # Reset to default token = task_id_var.set('') try: assert task_id_var.get() == '' finally: task_id_var.reset(token) def test_task_id_var_set_and_get(self): """Test setting and getting task_id_var.""" token = task_id_var.set('my-task-123') try: assert task_id_var.get() == 'my-task-123' finally: task_id_var.reset(token) def test_task_id_var_reset(self): """Test resetting task_id_var to previous value.""" # Set initial value token1 = task_id_var.set('task-1') assert task_id_var.get() == 'task-1' # Set new value token2 = task_id_var.set('task-2') assert task_id_var.get() == 'task-2' # Reset to previous task_id_var.reset(token2) assert task_id_var.get() == 'task-1' # Reset to original task_id_var.reset(token1) class TestLoggerLevel: """Tests for LOGGER_LEVEL configuration.""" @patch.dict('os.environ', {'LOGGER_LEVEL': 'DEBUG'}) def test_logger_level_from_environment(self): """Test that LOGGER_LEVEL reads from environment.""" # Need to reload module to pick up env var import importlib from abacus_common_logic.utils import logging as logging_module importlib.reload(logging_module) assert logging_module.LOGGER_LEVEL == logging.DEBUG