# Tests for logger module import pytest import json import logging from unittest.mock import Mock, patch from io import StringIO from src.utils.logger import ( StructuredFormatter, SimpleFormatter, setup_logger, get_logger, log_query_execution, log_database_operation, log_agent_step, log_query_rewrite, log_conversation_summary, log_context_filtering, log_sql_generation, log_database_query_execution, filter_sensitive_data, aggregate_logs, generate_request_id, get_request_id, clear_request_id, _request_id_local, ) class TestStructuredFormatter: """Tests for StructuredFormatter class.""" def test_formatter_creation(self): """Test basic formatter creation.""" formatter = StructuredFormatter() assert formatter._include_timestamp is True def test_formatter_without_timestamp(self): """Test formatter without timestamp.""" formatter = StructuredFormatter(include_timestamp=False) assert formatter._include_timestamp is False def test_format_basic_record(self): """Test formatting basic log record.""" formatter = StructuredFormatter(include_timestamp=False) record = logging.LogRecord( name="test_logger", level=logging.INFO, pathname="test.py", lineno=10, msg="Test message", args=(), exc_info=None, ) output = formatter.format(record) data = json.loads(output) assert data["level"] == "INFO" assert data["message"] == "Test message" assert data["logger"] == "test_logger" assert data["line"] == 10 def test_format_includes_timestamp(self): """Test formatting includes timestamp when enabled.""" formatter = StructuredFormatter(include_timestamp=True) record = logging.LogRecord( name="test", level=logging.INFO, pathname="test.py", lineno=1, msg="Test", args=(), exc_info=None, ) output = formatter.format(record) data = json.loads(output) assert "timestamp" in data def test_format_with_query_extra(self): """Test formatting with query in extra data.""" formatter = StructuredFormatter(include_timestamp=False) record = logging.LogRecord( name="test", level=logging.INFO, pathname="test.py", lineno=1, msg="Query executed", args=(), exc_info=None, ) record.query = "SELECT * FROM users" output = formatter.format(record) data = json.loads(output) assert data["query"] == "SELECT * FROM users" def test_format_with_execution_time(self): """Test formatting with execution time.""" formatter = StructuredFormatter(include_timestamp=False) record = logging.LogRecord( name="test", level=logging.INFO, pathname="test.py", lineno=1, msg="Query done", args=(), exc_info=None, ) record.execution_time = 1.234 output = formatter.format(record) data = json.loads(output) assert data["execution_time"] == 1.234 def test_format_with_extra_data(self): """Test formatting with extra data dictionary.""" formatter = StructuredFormatter(include_timestamp=False) record = logging.LogRecord( name="test", level=logging.INFO, pathname="test.py", lineno=1, msg="Message", args=(), exc_info=None, ) record.extra_data = {"key1": "value1", "key2": 123} output = formatter.format(record) data = json.loads(output) assert data["data"]["key1"] == "value1" assert data["data"]["key2"] == 123 def test_format_with_exception(self): """Test formatting with exception info.""" formatter = StructuredFormatter(include_timestamp=False) try: raise ValueError("Test error") except ValueError: import sys exc_info = sys.exc_info() record = logging.LogRecord( name="test", level=logging.ERROR, pathname="test.py", lineno=1, msg="Error occurred", args=(), exc_info=exc_info, ) output = formatter.format(record) data = json.loads(output) assert "exception" in data assert "ValueError" in data["exception"] class TestSimpleFormatter: """Tests for SimpleFormatter class.""" def test_formatter_creation(self): """Test simple formatter creation.""" formatter = SimpleFormatter() assert formatter._fmt is not None def test_formatter_custom_format(self): """Test simple formatter with custom format.""" custom_fmt = "%(levelname)s: %(message)s" formatter = SimpleFormatter(fmt=custom_fmt) record = logging.LogRecord( name="test", level=logging.INFO, pathname="test.py", lineno=1, msg="Test message", args=(), exc_info=None, ) output = formatter.format(record) assert "INFO: Test message" in output class TestSetupLogger: """Tests for setup_logger function.""" def test_setup_basic_logger(self): """Test setting up basic logger.""" logger = setup_logger("test_basic") assert logger.name == "test_basic" assert logger.level == logging.INFO assert len(logger.handlers) >= 1 for handler in logger.handlers: logger.removeHandler(handler) def test_setup_logger_with_level(self): """Test setting up logger with custom level.""" logger = setup_logger("test_level", level=logging.DEBUG) assert logger.level == logging.DEBUG for handler in logger.handlers: logger.removeHandler(handler) def test_setup_structured_logger(self): """Test setting up structured logger.""" logger = setup_logger("test_structured", structured=True) assert any( isinstance(h.formatter, StructuredFormatter) for h in logger.handlers ) for handler in logger.handlers: logger.removeHandler(handler) def test_setup_logger_idempotent(self): """Test that setup_logger is idempotent.""" logger1 = setup_logger("test_idempotent") handler_count = len(logger1.handlers) logger2 = setup_logger("test_idempotent") assert logger1 is logger2 assert len(logger2.handlers) == handler_count for handler in logger1.handlers: logger1.removeHandler(handler) class TestGetLogger: """Tests for get_logger function.""" def test_get_logger_returns_logger(self): """Test get_logger returns a logger.""" logger = get_logger("test_get") assert isinstance(logger, logging.Logger) assert logger.name == "test_get" def test_get_logger_same_name_same_instance(self): """Test get_logger returns same instance for same name.""" logger1 = get_logger("test_same") logger2 = get_logger("test_same") assert logger1 is logger2 class TestLogQueryExecution: """Tests for log_query_execution function.""" def test_log_successful_query(self): """Test logging successful query.""" logger = Mock() log_query_execution( logger=logger, query="SELECT * FROM users", execution_time=0.5, mode="database", success=True, ) logger.info.assert_called_once() call_args = logger.info.call_args assert "database" in call_args[0][0] assert "0.500" in call_args[0][0] def test_log_failed_query(self): """Test logging failed query.""" logger = Mock() log_query_execution( logger=logger, query="SELECT * FROM users", execution_time=1.0, mode="database", success=False, error="Connection failed", ) logger.error.assert_called_once() call_args = logger.error.call_args assert "Connection failed" in call_args[0][0] def test_log_query_with_extra_data(self): """Test logging query includes extra data.""" logger = Mock() log_query_execution( logger=logger, query="test query", execution_time=0.1, mode="documents", success=True, ) call_args = logger.info.call_args extra = call_args[1]["extra"] assert extra["query"] == "test query" assert extra["execution_time"] == 0.1 class TestLogDatabaseOperation: """Tests for log_database_operation function.""" def test_log_successful_operation(self): """Test logging successful database operation.""" logger = Mock() log_database_operation( logger=logger, operation="SELECT", table="users", execution_time=0.2, row_count=10, success=True, ) logger.info.assert_called_once() call_args = logger.info.call_args assert "SELECT" in call_args[0][0] assert "users" in call_args[0][0] def test_log_failed_operation(self): """Test logging failed database operation.""" logger = Mock() log_database_operation( logger=logger, operation="INSERT", table="orders", success=False, error="Constraint violation", ) logger.error.assert_called_once() call_args = logger.error.call_args assert "Constraint violation" in call_args[0][0] def test_log_operation_minimal_info(self): """Test logging operation with minimal info.""" logger = Mock() log_database_operation( logger=logger, operation="COMMIT", ) logger.info.assert_called_once() class TestLogAgentStep: """Tests for log_agent_step function.""" def test_log_successful_step(self): """Test logging successful agent step.""" logger = Mock() log_agent_step( logger=logger, step_name="retrieve_documents", step_type="tool", tool_name="document_retriever", execution_time=0.5, success=True, ) logger.info.assert_called_once() call_args = logger.info.call_args assert "retrieve_documents" in call_args[0][0] assert "tool" in call_args[0][0] def test_log_failed_step(self): """Test logging failed agent step.""" logger = Mock() log_agent_step( logger=logger, step_name="generate_sql", step_type="tool", tool_name="database_query", success=False, ) logger.warning.assert_called_once() call_args = logger.warning.call_args assert "failed" in call_args[0][0] def test_log_step_with_result_summary(self): """Test logging step with result summary.""" logger = Mock() log_agent_step( logger=logger, step_name="classify_intent", step_type="classification", success=True, result_summary="database_query", ) logger.info.assert_called_once() call_args = logger.info.call_args assert "database_query" in call_args[0][0] class TestRequestTracing: def setup_method(self): clear_request_id() def teardown_method(self): clear_request_id() def test_generate_request_id_returns_uuid(self): rid = generate_request_id() assert rid is not None assert len(rid) == 36 assert rid.count("-") == 4 def test_get_request_id_returns_generated(self): rid = generate_request_id() assert get_request_id() == rid def test_get_request_id_none_when_not_set(self): assert get_request_id() is None def test_clear_request_id(self): generate_request_id() clear_request_id() assert get_request_id() is None def test_generate_unique_ids(self): rid1 = generate_request_id() rid2 = generate_request_id() assert rid1 != rid2 def test_request_id_included_in_log_data(self): generate_request_id() logger = Mock() log_query_rewrite(logger, "original", "rewritten", "llm", 10.0, False) data = logger.info.call_args[1]["extra"]["extra_data"] assert "request_id" in data assert data["request_id"] == get_request_id() class TestLogQueryRewrite: def setup_method(self): clear_request_id() @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_logs_successful_rewrite(self): logger = Mock() log_query_rewrite(logger, "Bu nedir?", "Atlas nedir?", "llm", 15.5, False) logger.info.assert_called_once() msg = logger.info.call_args[0][0] assert "method=llm" in msg assert "cache=False" in msg assert "changed=True" in msg @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_logs_unchanged_query(self): logger = Mock() log_query_rewrite(logger, "merhaba", "merhaba", "none", 0.5, False) msg = logger.info.call_args[0][0] assert "changed=False" in msg @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_logs_cache_hit(self): logger = Mock() log_query_rewrite(logger, "q", "r", "cache", 0.1, True) msg = logger.info.call_args[0][0] assert "cache=True" in msg @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_truncates_long_queries(self): logger = Mock() long_q = "x" * 200 log_query_rewrite(logger, long_q, long_q, "llm", 5.0, False) data = logger.info.call_args[1]["extra"]["extra_data"] assert len(data["original_query"]) <= 100 assert len(data["rewritten_query"]) <= 100 @patch("src.utils.logger.LOG_SAMPLING_RATE", 0.0) def test_sampling_skips_log(self): logger = Mock() log_query_rewrite(logger, "q", "r", "llm", 5.0, False) logger.info.assert_not_called() @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_includes_hashes(self): logger = Mock() log_query_rewrite(logger, "original", "rewritten", "llm", 5.0, False) data = logger.info.call_args[1]["extra"]["extra_data"] assert "original_hash" in data assert "rewritten_hash" in data assert len(data["original_hash"]) == 8 class TestLogConversationSummary: def setup_method(self): clear_request_id() @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_logs_summary(self): logger = Mock() log_conversation_summary(logger, 20, 8, "Ozet metni burada", 500.0) logger.debug.assert_called_once() msg = logger.debug.call_args[0][0] assert "20->8" in msg @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_calculates_reduction_ratio(self): logger = Mock() log_conversation_summary(logger, 10, 3, "sum", 100.0) data = logger.debug.call_args[1]["extra"]["extra_data"] assert data["reduction_ratio"] == 0.7 @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_handles_zero_original(self): logger = Mock() log_conversation_summary(logger, 0, 0, "", 0.0) data = logger.debug.call_args[1]["extra"]["extra_data"] assert data["reduction_ratio"] == 0.0 @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_truncates_summary_preview(self): logger = Mock() long_summary = "z" * 500 log_conversation_summary(logger, 10, 3, long_summary, 50.0) data = logger.debug.call_args[1]["extra"]["extra_data"] assert len(data["summary_preview"]) <= 200 @patch("src.utils.logger.LOG_SAMPLING_RATE", 0.0) def test_sampling_skips(self): logger = Mock() log_conversation_summary(logger, 20, 5, "sum", 100.0) logger.debug.assert_not_called() class TestLogContextFiltering: def setup_method(self): clear_request_id() @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_logs_filtering(self): logger = Mock() log_context_filtering(logger, "database_query", 10, 4, "keyword_filter") logger.debug.assert_called_once() msg = logger.debug.call_args[0][0] assert "intent=database_query" in msg assert "10->4" in msg @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_includes_reduction(self): logger = Mock() log_context_filtering(logger, "document_query", 8, 8, "full_context") data = logger.debug.call_args[1]["extra"]["extra_data"] assert data["reduction"] == 0 @patch("src.utils.logger.LOG_SAMPLING_RATE", 0.0) def test_sampling_skips(self): logger = Mock() log_context_filtering(logger, "hybrid", 10, 5, "dual") logger.debug.assert_not_called() class TestLogSqlGeneration: def setup_method(self): clear_request_id() @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_logs_successful_generation(self): logger = Mock() log_sql_generation(logger, "kac satis var", "SELECT COUNT(*) FROM sales", 500, 4, 120.0) logger.info.assert_called_once() msg = logger.info.call_args[0][0] assert "schema=500B" in msg @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_logs_failed_generation(self): logger = Mock() log_sql_generation(logger, "query", "", 100, 2, 50.0, success=False, error="LLM timeout") logger.error.assert_called_once() msg = logger.error.call_args[0][0] assert "LLM timeout" in msg @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_errors_always_logged_regardless_of_sampling(self): logger = Mock() log_sql_generation(logger, "q", "", 100, 0, 10.0, success=False, error="fail") logger.error.assert_called_once() @patch("src.utils.logger.LOG_SAMPLING_RATE", 0.0) def test_sampling_skips_success_not_error(self): logger = Mock() log_sql_generation(logger, "q", "SELECT 1", 100, 0, 10.0, success=True) logger.info.assert_not_called() @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_truncates_long_sql(self): logger = Mock() long_sql = "SELECT " + "x" * 500 log_sql_generation(logger, "q", long_sql, 100, 0, 10.0) data = logger.info.call_args[1]["extra"]["extra_data"] assert len(data["generated_sql"]) <= 200 @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_includes_sql_hash(self): logger = Mock() log_sql_generation(logger, "q", "SELECT 1", 100, 0, 10.0) data = logger.info.call_args[1]["extra"]["extra_data"] assert "sql_hash" in data assert len(data["sql_hash"]) == 12 class TestLogDatabaseQueryExecution: def setup_method(self): clear_request_id() @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) @patch("src.utils.logger.LOG_SLOW_QUERY_THRESHOLD_MS", 1000) def test_logs_successful_fast_query(self): logger = Mock() log_database_query_execution(logger, "SELECT 1", 5, 50.0, True) logger.info.assert_called_once() assert "rows=5" in logger.info.call_args[0][0] @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) @patch("src.utils.logger.LOG_SLOW_QUERY_THRESHOLD_MS", 1000) def test_logs_slow_query_as_warning(self): logger = Mock() log_database_query_execution(logger, "SELECT * FROM big_table", 1000, 1500.0, True) logger.warning.assert_called_once() assert "Slow" in logger.warning.call_args[0][0] @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_logs_failed_query_as_error(self): logger = Mock() log_database_query_execution(logger, "SELECT bad", 0, 10.0, False, error="syntax error") logger.error.assert_called_once() assert "syntax error" in logger.error.call_args[0][0] @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_errors_always_logged(self): logger = Mock() log_database_query_execution(logger, "bad sql", 0, 5.0, False, error="err") logger.error.assert_called_once() @patch("src.utils.logger.LOG_SAMPLING_RATE", 0.0) @patch("src.utils.logger.LOG_SLOW_QUERY_THRESHOLD_MS", 1000) def test_sampling_skips_successful_fast(self): logger = Mock() log_database_query_execution(logger, "SELECT 1", 1, 10.0, True) logger.info.assert_not_called() logger.warning.assert_not_called() @patch("src.utils.logger.LOG_SAMPLING_RATE", 1.0) def test_includes_sql_hash(self): logger = Mock() log_database_query_execution(logger, "SELECT 1", 1, 5.0, True) data = logger.info.call_args[1]["extra"]["extra_data"] assert "sql_hash" in data class TestFilterSensitiveData: def test_redacts_email(self): data = {"message": "Contact user@example.com for info"} filtered = filter_sensitive_data(data) assert "user@example.com" not in filtered["message"] assert "[EMAIL_REDACTED]" in filtered["message"] def test_redacts_phone_number(self): data = {"info": "Call +90 555 123 4567 now"} filtered = filter_sensitive_data(data) assert "555 123 4567" not in filtered["info"] assert "[PHONE_REDACTED]" in filtered["info"] def test_redacts_api_token(self): data = {"auth": "Bearer sk-abc123def456"} filtered = filter_sensitive_data(data) assert "sk-abc123def456" not in filtered["auth"] assert "[TOKEN_REDACTED]" in filtered["auth"] def test_redacts_api_key(self): data = {"config": "api_key=secret123"} filtered = filter_sensitive_data(data) assert "secret123" not in filtered["config"] def test_preserves_non_sensitive_data(self): data = {"query": "SELECT * FROM users", "count": 42} filtered = filter_sensitive_data(data) assert filtered["query"] == "SELECT * FROM users" assert filtered["count"] == 42 def test_handles_nested_dicts(self): data = {"outer": {"email": "test@test.com"}} filtered = filter_sensitive_data(data) assert "[EMAIL_REDACTED]" in filtered["outer"]["email"] def test_handles_empty_dict(self): assert filter_sensitive_data({}) == {} def test_preserves_non_string_values(self): data = {"count": 10, "rate": 0.95, "active": True, "items": None} filtered = filter_sensitive_data(data) assert filtered == data class TestAggregateLogs: def test_counts_operations(self): entries = [ {"operation": "query_rewrite", "time_ms": 10}, {"operation": "query_rewrite", "time_ms": 20}, {"operation": "sql_generation", "generation_time_ms": 100}, ] result = aggregate_logs(entries) assert result["total_operations"] == 3 assert result["by_operation_type"]["query_rewrite"]["count"] == 2 assert result["by_operation_type"]["sql_generation"]["count"] == 1 def test_calculates_avg_time(self): entries = [ {"operation": "query_rewrite", "time_ms": 10}, {"operation": "query_rewrite", "time_ms": 30}, ] result = aggregate_logs(entries) assert result["by_operation_type"]["query_rewrite"]["avg_time_ms"] == 20.0 def test_collects_errors(self): entries = [ {"operation": "sql_generation", "success": False, "error": "timeout"}, {"operation": "sql_generation", "success": True}, ] result = aggregate_logs(entries) assert len(result["errors"]) == 1 assert result["by_operation_type"]["sql_generation"]["success_rate"] == 0.5 @patch("src.utils.logger.LOG_SLOW_QUERY_THRESHOLD_MS", 100) def test_collects_slow_operations(self): entries = [ {"operation": "db_exec", "execution_time_ms": 200}, {"operation": "db_exec", "execution_time_ms": 50}, ] result = aggregate_logs(entries) assert len(result["slow_operations"]) == 1 def test_handles_empty_entries(self): result = aggregate_logs([]) assert result["total_operations"] == 0 assert result["by_operation_type"] == {} def test_includes_time_range(self): result = aggregate_logs([], time_range_hours=48) assert result["time_range_hours"] == 48 def test_limits_error_list_size(self): entries = [{"operation": "op", "error": f"err{i}"} for i in range(100)] result = aggregate_logs(entries) assert len(result["errors"]) <= 50 def test_unknown_operation_handled(self): entries = [{"time_ms": 5}] result = aggregate_logs(entries) assert "unknown" in result["by_operation_type"]