acb / tests /test_logger.py
Kagan Tek
merge
79d4fd5
Raw
History Blame Contribute Delete
26.1 kB
# 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"]