diff --git a/CHANGELOG.md b/CHANGELOG.md index 9c51b93..4950ee4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -18,6 +18,7 @@ and versions are tracked in the repo-root `VERSION` file. ### Changed +- Reuse secure log lock descriptors and cache source paths per invocation; logging I/O failures stay inside logging (#381). - Align the Typer support floor with the tested matrix and cover representative minimum/maximum Typer and Click version pairings. - Prefix formula-leading CSV/TSV cells with an apostrophe by default to protect diff --git a/docs/performance.md b/docs/performance.md index 96e0797..5bb71f8 100644 --- a/docs/performance.md +++ b/docs/performance.md @@ -129,3 +129,11 @@ extension discovery caches, and run-bundle retention. Ctrl+C is tested through both the lifecycle boundary and a real POSIX subprocess signal. Windows keeps the portable lifecycle and persistence checks while skipping only assertions that require POSIX signal or descriptor semantics. + +### Logging hot path + +Secure log handlers keep their private sidecar descriptor open until cleanup, +with an advisory lock around each append and a fresh descriptor after fork. +Human formatters cache up to 256 source paths for the current invocation and +project binding. Repeated paths require no filesystem resolution. Sidecar I/O +errors are routed through `logging.Handler.handleError` and do not fail commands. diff --git a/lib/python/base_cli/logging.py b/lib/python/base_cli/logging.py index 30309ad..6f66bdb 100644 --- a/lib/python/base_cli/logging.py +++ b/lib/python/base_cli/logging.py @@ -167,6 +167,9 @@ def secure_log_file_permissions(log_file: Path) -> None: class SecureLogFileHandler(logging.FileHandler): def __init__(self, filename: str | os.PathLike[str], *args: object, **kwargs: object) -> None: self._lock_path = Path(filename).with_name(f".{Path(filename).name}.lock") + self._lock_stream: BinaryIO | None = None + self._lock_identity: tuple[int, int] | None = None + self._lock_pid = os.getpid() super().__init__(filename, *args, **kwargs) # type: ignore[arg-type] def _open(self) -> TextIOWrapper: @@ -184,22 +187,79 @@ def _open(self) -> TextIOWrapper: raise def emit(self, record: logging.LogRecord) -> None: - lock_stream = _open_log_lock(self._lock_path) try: - _lock_log_stream(lock_stream) - super().emit(record) + # An inherited flock descriptor would share ownership with the parent. + if self._lock_pid != os.getpid(): + inherited = self._lock_stream + self._lock_stream = None + self._lock_identity = None + self._lock_pid = os.getpid() + if inherited is not None: + try: + inherited.close() + except OSError: + pass + if self._lock_stream is not None: + try: + current = os.stat(self._lock_path, follow_symlinks=False) + stream_stat = os.fstat(self._lock_stream.fileno()) + except FileNotFoundError: + current = None + stream_stat = None + if ( + current is None + or stream_stat is None + or ( + self._lock_identity != (current.st_dev, current.st_ino) + or (stream_stat.st_dev, stream_stat.st_ino) != self._lock_identity + ) + ): + stale = self._lock_stream + self._lock_stream = None + self._lock_identity = None + stale.close() + else: + try: + restrict_file(self._lock_path) + except FileNotFoundError: + stale = self._lock_stream + self._lock_stream = None + self._lock_identity = None + stale.close() + if self._lock_stream is None: + self._lock_stream = _open_log_lock(self._lock_path) + lock_stat = os.fstat(self._lock_stream.fileno()) + self._lock_identity = (lock_stat.st_dev, lock_stat.st_ino) + _lock_log_stream(self._lock_stream) + try: + super().emit(record) + finally: + _unlock_log_stream(self._lock_stream) + except RecursionError: + raise + except Exception: + self.handleError(record) + + def close(self) -> None: + self.acquire() + try: + stream, self._lock_stream = self._lock_stream, None + self._lock_identity = None + try: + if stream is not None: + stream.close() + finally: + super().close() finally: - _unlock_log_stream(lock_stream) - lock_stream.close() + self.release() def _open_log_lock(path: Path) -> BinaryIO: path.parent.mkdir(parents=True, exist_ok=True) stream = path.open("a+b") try: - if stream.seek(0, os.SEEK_END) == 0: - stream.write(b"0") - stream.flush() + # Byte-range locks may extend beyond EOF. Writing a sentinel before + # acquiring the lock races a Windows writer already holding byte zero. restrict_file(path) return stream except BaseException: @@ -253,13 +313,19 @@ def __init__(self, *, use_utc: bool | None = None, use_color: bool = False) -> N resolved_use_utc = legacy_use_utc == "1" self.use_utc = resolved_use_utc self.use_color = use_color + self._source_key: tuple[object, ...] | None = None + self._source_roots: tuple[Path, ...] = () + self._source_cache: dict[str, str] = {} + self._cache_lock = RLock() + self._time_key: tuple[object, ...] | None = None + self._time_text = "" datefmt = "%Y-%m-%d %H:%M:%S UTC" if self.use_utc else "%Y-%m-%d %H:%M:%S %z" super().__init__(datefmt=datefmt) self.converter = time.gmtime if self.use_utc else time.localtime def format(self, record: logging.LogRecord) -> str: timestamp = self.formatTime(record, self.datefmt) - source = _source_path(record) + source = self._source_path(record) level = _level_name(record) line = f"{timestamp} {level:<7} {source}:{record.lineno} {record.getMessage()}" if record.exc_info: @@ -273,6 +339,49 @@ def format(self, record: logging.LogRecord) -> str: color = _LEVEL_COLORS.get(record.levelno) return f"{color}{line}{_COLOR_RESET}" if color else line + def formatTime(self, record: logging.LogRecord, datefmt: str | None = None) -> str: + with self._cache_lock: + if datefmt is None: + return super().formatTime(record, datefmt) + # Human timestamps have second precision. Repeated calls to localtime + # and strftime otherwise re-read timezone state on some platforms. + key = (record.created // 1, datefmt, self.converter, os.environ.get("TZ"), time.tzname) + if key != self._time_key: + self._time_text = super().formatTime(record, datefmt) + self._time_key = key + return self._time_text + + def _source_path(self, record: logging.LogRecord) -> str: + with self._cache_lock: + cwd = current_working_dir() + try: + context = get_current_context() + except RuntimeError: + key: tuple[object, ...] = (None, cwd) + roots: tuple[Path, ...] = (cwd,) + else: + key = (context.run_id, context.application_home, context.project_root, cwd) + roots = tuple(p for p in (context.application_home, context.project_root) if p is not None) + if key != self._source_key: + self._source_roots = tuple(p.resolve() for p in (*roots, cwd)) + self._source_key = key + self._source_cache.clear() + cached = self._source_cache.get(record.pathname) + if cached is not None: + return cached + path = Path(record.pathname).resolve() + source = str(path) + for root in self._source_roots: + try: + source = str(path.relative_to(root)) + break + except ValueError: + continue + if len(self._source_cache) >= 256: + self._source_cache.clear() + self._source_cache[record.pathname] = source + return source + def _level_name(record: logging.LogRecord) -> str: if record.levelno == logging.WARNING: @@ -282,41 +391,6 @@ def _level_name(record: logging.LogRecord) -> str: return record.levelname -def _source_path(record: logging.LogRecord) -> str: - path = Path(record.pathname) - candidates = [] - application_home = _active_application_home() - if application_home is not None: - candidates.append(application_home) - project_root = _active_project_root() - if project_root is not None: - candidates.append(project_root) - candidates.append(current_working_dir()) - - for root in candidates: - try: - return str(path.resolve().relative_to(root.resolve())) - except ValueError: - continue - return str(path.resolve()) - - -def _active_project_root() -> Path | None: - try: - context = get_current_context() - except RuntimeError: - return None - return context.project_root - - -def _active_application_home() -> Path | None: - try: - context = get_current_context() - except RuntimeError: - return None - return context.application_home - - def log_invocation( logger: logging.Logger, argv: list[str], diff --git a/tests/test_logging_hot_path.py b/tests/test_logging_hot_path.py new file mode 100644 index 0000000..df3c43d --- /dev/null +++ b/tests/test_logging_hot_path.py @@ -0,0 +1,139 @@ +from __future__ import annotations + +import io +import logging +import os +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from unittest.mock import patch + +import base_cli +import base_cli.logging as module +import pytest +from base_cli.testing import invoke + + +def test_sidecar_is_opened_once_and_closed(tmp_path: Path) -> None: + handler = module.SecureLogFileHandler(tmp_path / "run.log") + record = logging.LogRecord("test", logging.INFO, __file__, 1, "message", (), None) + with patch.object(module, "_open_log_lock", wraps=module._open_log_lock) as opened: + handler.emit(record) + handler.emit(record) + assert opened.call_count == 1 + stream = handler._lock_stream + handler.close() + assert stream.closed + + +@pytest.mark.skipif(os.name == "nt", reason="Windows does not permit unlinking an open sidecar") +def test_deleted_sidecar_reopens_and_preserves_later_records(tmp_path: Path) -> None: + handler = module.SecureLogFileHandler(tmp_path / "run.log") + try: + handler.emit(logging.LogRecord("test", logging.INFO, __file__, 1, "before", (), None)) + handler._lock_path.unlink() + for message in ("after-0", "after-1", "after-2"): + handler.emit(logging.LogRecord("test", logging.INFO, __file__, 1, message, (), None)) + finally: + handler.close() + + log_text = (tmp_path / "run.log").read_text(encoding="utf-8") + assert log_text.splitlines() == ["before", "after-0", "after-1", "after-2"] + + +def test_shared_formatter_handles_concurrent_user_and_file_logging(tmp_path: Path) -> None: + user_stream = io.StringIO() + formatter = module.CliFormatter() + logger = base_cli.configure_logger( + "shared-formatter-race", + tmp_path / "run.log", + debug=True, + stream=user_stream, + formatter=formatter, + propagate=False, + ) + try: + with ThreadPoolExecutor(max_workers=8) as executor: + list(executor.map(logger.info, (f"message-{index}" for index in range(64)))) + assert user_stream.getvalue().count("message-") == 64 + assert (tmp_path / "run.log").read_text(encoding="utf-8").count("message-") == 64 + finally: + for handler in list(logger.handlers): + handler.close() + logger.removeHandler(handler) + + +def test_forked_handler_recovers_after_inherited_lock_close_failure(tmp_path: Path) -> None: + handler = module.SecureLogFileHandler(tmp_path / "run.log") + record = logging.LogRecord("test", logging.INFO, __file__, 1, "message", (), None) + + class FailingStream: + def close(self) -> None: + raise OSError("already closed") + + handler._lock_stream = FailingStream() # type: ignore[assignment] + handler._lock_pid = os.getpid() - 1 + try: + handler.emit(record) + assert handler._lock_pid == os.getpid() + assert handler._lock_stream is not None + finally: + handler.close() + + +def test_recursion_errors_follow_stdlib_handler_contract(tmp_path: Path) -> None: + handler = module.SecureLogFileHandler(tmp_path / "run.log") + record = logging.LogRecord("test", logging.INFO, __file__, 1, "message", (), None) + try: + with patch.object(logging.FileHandler, "emit", side_effect=RecursionError("recursive")): + with pytest.raises(RecursionError, match="recursive"): + handler.emit(record) + finally: + handler.close() + + +def test_logging_lock_failures_do_not_fail_command(tmp_path: Path) -> None: + app = base_cli.App(name="logging-io-failure") + + @app.command() + def main(ctx: base_cli.Context) -> None: + for failure in (PermissionError("unwritable log directory"), OSError("volume full")): + with patch.object(module, "_lock_log_stream", side_effect=failure): + ctx.log.info("still completes") + ctx.log.info("recovers") + + result = invoke(app, [], home=tmp_path) + assert result.exit_code == 0 + assert "Logging error" in result.stderr + + +def test_formatter_repeated_paths_do_not_resolve_again(tmp_path: Path) -> None: + app = base_cli.App(name="cached-log-source") + + @app.command() + def main(ctx: base_cli.Context) -> None: + ctx.log.info("warm cache") + with patch.object(Path, "resolve", side_effect=AssertionError("unexpected resolution")): + ctx.log.info("cached source") + + assert invoke(app, [], home=tmp_path).exit_code == 0 + + +def test_timestamp_cache_preserves_seconds_and_timezone_format() -> None: + for use_utc in (False, True): + formatter = module.CliFormatter(use_utc=use_utc) + reference = logging.Formatter(datefmt=formatter.datefmt) + reference.converter = formatter.converter + record = logging.LogRecord("test", logging.INFO, __file__, 1, "message", (), None) + for created in (1000.1, 1000.9, 1001.0, 1002.3): + record.created = created + assert formatter.formatTime(record, formatter.datefmt) == reference.formatTime(record, formatter.datefmt) + + +def test_opening_lock_does_not_write_an_unlocked_sentinel(tmp_path: Path) -> None: + path = tmp_path / "append.lock" + with module._open_log_lock(path) as stream: + module._lock_log_stream(stream) + try: + assert path.stat().st_size == 0 + finally: + module._unlock_log_stream(stream)