diff --git a/tests/conftest.py b/tests/conftest.py index 698c782..c1b49b3 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,3 +1,5 @@ +import logging + import pytest from django.tasks import default_task_backend @@ -9,6 +11,17 @@ def _flush_keys(backend) -> None: backend.client.delete(*keys) +@pytest.fixture(autouse=True) +def restore_root_logger(): + """Restore root logger handlers and level after each test.""" + root_logger = logging.getLogger() + handlers = root_logger.handlers[:] + level = root_logger.level + yield + root_logger.handlers[:] = handlers + root_logger.setLevel(level) + + @pytest.fixture(autouse=True) def flush_default_backend(): """Flush all threadmill keys and reset async client before and after each test.""" diff --git a/tests/test_executor.py b/tests/test_executor.py index c6bb440..19ecbbf 100644 --- a/tests/test_executor.py +++ b/tests/test_executor.py @@ -26,6 +26,7 @@ boom_with_retry, count_users, echo, + log_message, ) from threadmill.backends.base import Broker from threadmill.executor import ( @@ -33,6 +34,7 @@ TaskExecutor, WorkerProcess, WorkerThread, + configure_logging, handler, ) @@ -158,6 +160,44 @@ def test_format__includes_exception_traceback(self) -> None: assert "ValueError: boom" in payload["exception"] +class TestConfigureLogging: + """Tests for the configure_logging function.""" + + @pytest.fixture(autouse=True) + def restore_log_formatter(self): + """Restore the shared log formatter after each test.""" + formatter = handler.formatter + yield + handler.setFormatter(formatter) + + def test_configure_logging__installs_handler_on_root(self): + """Route the records of every logger through the shared handler.""" + formatter = logging.Formatter("%(levelname)s %(message)s") + configure_logging(formatter) + root_logger = logging.getLogger() + assert root_logger.handlers == [handler] + assert root_logger.level == logging.INFO + assert handler.formatter is formatter + + def test_configure_logging__replaces_foreign_handlers(self): + """Replace handlers of existing loggers so their records reach the root.""" + task_logger = logging.getLogger("tests.testapp.tasks") + task_logger.addHandler(logging.NullHandler()) + task_logger.propagate = False + configure_logging(JsonFormatter()) + assert task_logger.handlers == [] + assert task_logger.propagate is True + + def test_configure_logging__keeps_placeholder_loggers(self): + """Ignore placeholder entries in the logger registry.""" + logging.getLogger("tests.test_executor.placeholder.child") + configure_logging(JsonFormatter()) + assert isinstance( + logging.root.manager.loggerDict["tests.test_executor.placeholder"], + logging.PlaceHolder, + ) + + class TestTaskExecutor: """Tests for the TaskExecutor dataclass and its methods.""" @@ -247,6 +287,43 @@ def test_run__processes_enqueued_tasks_end_to_end(self): assert {r.id for r in results} == {r.id for r in enqueued} assert all(r.status == TaskResultStatus.SUCCESSFUL for r in results) + def test_run__routes_task_logs_to_stdout(self, capfd): + """Emit task log records as JSON on standard output.""" + enqueued = default_task_backend.enqueue(log_message, args=["hello from task"]) + original_start_method = multiprocessing.get_start_method() + # A forkserver worker inherits the stdout of the long-lived forkserver + # instead of the file descriptor this fixture replaces, so its records + # would never reach capfd. + multiprocessing.set_start_method("spawn", force=True) + try: + TaskExecutor( + backend=default_task_backend, + workers=1, + threads=1, + queues=("default",), + exit_empty=True, + ).run() + finally: + multiprocessing.set_start_method(original_start_method, force=True) + + captured = capfd.readouterr() + assert "hello from task" not in captured.err + records = [ + json.loads(line) + for line in captured.out.splitlines() + if line.startswith("{") + ] + assert any( + record["logger"] == "tests.testapp.tasks" + and record["message"] == "hello from task" + and record["level"] == "INFO" + for record in records + ) + assert ( + default_task_backend.get_result(enqueued.id).status + is TaskResultStatus.SUCCESSFUL + ) + def test_run__executes_model_task_in_spawned_worker(self): """run() executes a model-accessing task in a spawned worker process.""" original_start_method = multiprocessing.get_start_method() diff --git a/tests/testapp/tasks.py b/tests/testapp/tasks.py index 51a4286..e8481cb 100644 --- a/tests/testapp/tasks.py +++ b/tests/testapp/tasks.py @@ -22,6 +22,13 @@ def boom(): raise ValueError("boom") +@task() +def log_message(message): + """Log the given message at INFO level (tests task log routing).""" + logger.info(message) + return message + + @task() def count_users(): """Count all users in the database (tests model access in workers).""" diff --git a/threadmill/executor.py b/threadmill/executor.py index 14df8b2..1f4a58c 100644 --- a/threadmill/executor.py +++ b/threadmill/executor.py @@ -8,6 +8,7 @@ import multiprocessing import random import socket +import sys import threading import time import typing @@ -81,10 +82,21 @@ def format(self, record: logging.LogRecord) -> str: logger = multiprocessing.get_logger() -handler = logging.StreamHandler() +handler = logging.StreamHandler(sys.stdout) handler.setFormatter(JsonFormatter()) -logger.addHandler(handler) -logger.setLevel(logging.INFO) + + +def configure_logging(formatter: logging.Formatter) -> None: + """Route every log record of this process through the threadmill handler.""" + handler.setFormatter(formatter) + for existing_logger in logging.root.manager.loggerDict.values(): + if isinstance(existing_logger, logging.Logger): + existing_logger.handlers.clear() + existing_logger.propagate = True + root_logger = logging.getLogger() + root_logger.handlers.clear() + root_logger.addHandler(handler) + root_logger.setLevel(logging.INFO) @dataclasses.dataclass(kw_only=True, slots=True) @@ -138,7 +150,7 @@ def create_worker_process(self) -> WorkerProcess: def run(self) -> None: """Start consuming tasks until shutdown is requested.""" - handler.setFormatter(self.log_formatter) + configure_logging(self.log_formatter) self.worker_processes = [ self.create_worker_process() for _ in range(self.process_count) ] @@ -213,8 +225,8 @@ def __init__( def run(self) -> None: """Start consumer execution inside this process.""" - handler.setFormatter(self.log_formatter) django.setup() + configure_logging(self.log_formatter) logger.info("Starting worker process %s", self.name) self.lock = threading.Lock() self.expired = threading.Event() diff --git a/threadmill/management/commands/threadmill.py b/threadmill/management/commands/threadmill.py index 17cad16..2c39524 100644 --- a/threadmill/management/commands/threadmill.py +++ b/threadmill/management/commands/threadmill.py @@ -88,8 +88,8 @@ def add_arguments(self, parser): parser.add_argument( "--log-format", help=( - "Logging format string for worker log records, e.g." - " '%%(levelname)s %%(message)s'. Defaults to JSON." + "Logging format string for all log records of the worker process," + " e.g. '%%(levelname)s %%(message)s'. Defaults to JSON." ), )