Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import logging

import pytest
from django.tasks import default_task_backend

Expand All @@ -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."""
Expand Down
77 changes: 77 additions & 0 deletions tests/test_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,13 +26,15 @@
boom_with_retry,
count_users,
echo,
log_message,
)
from threadmill.backends.base import Broker
from threadmill.executor import (
JsonFormatter,
TaskExecutor,
WorkerProcess,
WorkerThread,
configure_logging,
handler,
)

Expand Down Expand Up @@ -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."""

Expand Down Expand Up @@ -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()
Expand Down
7 changes: 7 additions & 0 deletions tests/testapp/tasks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)."""
Expand Down
22 changes: 17 additions & 5 deletions threadmill/executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
import multiprocessing
import random
import socket
import sys
import threading
import time
import typing
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
]
Expand Down Expand Up @@ -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()
Expand Down
4 changes: 2 additions & 2 deletions threadmill/management/commands/threadmill.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."
),
)

Expand Down