diff --git a/_python_utils_tests/clock.py b/_python_utils_tests/clock.py new file mode 100644 index 0000000..d0c5cbe --- /dev/null +++ b/_python_utils_tests/clock.py @@ -0,0 +1,56 @@ +"""A fake clock for the tests and doctests of ``python_utils.time``.""" + +import typing + +import pytest + +import python_utils.time + +#: The doctest that sleeps on the clock, by the name pytest gives it. +TIMEOUT_GENERATOR_DOCTEST: str = 'python_utils.time.timeout_generator' + + +class FakeClock: + """ + A clock that only moves when something sleeps on it. + + ``time.sleep`` promises to sleep at least as long as requested. A busy + machine sleeps tens of milliseconds longer, and that changes how many + items ``timeout_generator`` yields before its timeout. On this clock the + number of items depends on the arguments alone. + + Attributes: + now (float): The current time in seconds. + sleeps (list[float]): Every requested sleep, in order. + """ + + def __init__(self) -> None: + """Start at zero without any recorded sleeps.""" + self.now: float = 0.0 + self.sleeps: list[float] = [] + + def perf_counter(self) -> float: + """Return the current time, like ``time.perf_counter``.""" + return self.now + + def sleep(self, seconds: float) -> None: + """Record the sleep and move the clock forward, without waiting.""" + self.sleeps.append(seconds) + self.now += seconds + + +@pytest.fixture +def fake_clock(monkeypatch: pytest.MonkeyPatch) -> FakeClock: + """Replace the ``time`` module inside ``python_utils.time``.""" + clock: FakeClock = FakeClock() + monkeypatch.setattr(python_utils.time, 'time', clock) + return clock + + +@pytest.fixture(autouse=True) +def fake_clock_in_doctest(request: pytest.FixtureRequest) -> None: + """Run the ``timeout_generator`` doctest on the fake clock.""" + # pytest leaves `FixtureRequest.node` without a type. + node: pytest.Item = typing.cast(pytest.Item, request.node) + if node.name == TIMEOUT_GENERATOR_DOCTEST: + request.getfixturevalue('fake_clock') diff --git a/_python_utils_tests/test_converters.py b/_python_utils_tests/test_converters.py new file mode 100644 index 0000000..9ad4dc9 --- /dev/null +++ b/_python_utils_tests/test_converters.py @@ -0,0 +1,39 @@ +"""Tests for the conversion helpers in ``python_utils.converters``.""" + +import re + +import pytest + +from python_utils import converters + + +@pytest.mark.parametrize('regexp', [r'\d+', re.compile(r'\d+')]) +def test_to_int_regexp_without_group(regexp: re.Pattern[str] | str) -> None: + """Use the whole match when the pattern has no capture group.""" + assert converters.to_int('abc123', regexp=regexp) == 123 + + +@pytest.mark.parametrize('regexp', [r'\d+\.\d+', re.compile(r'\d+\.\d+')]) +def test_to_float_regexp_without_group( + regexp: re.Pattern[str] | str, +) -> None: + """Use the whole match when the pattern has no capture group.""" + assert converters.to_float('abc1.5', regexp=regexp) == 1.5 + + +def test_regexp_with_groups_keeps_its_group() -> None: + """Keep the last group for ``to_int`` and the first for ``to_float``.""" + assert converters.to_int('a1b2', regexp=r'(\d)\D(\d)') == 2 + assert converters.to_float('a1b2', regexp=r'(\d)\D(\d)') == 1.0 + + +@pytest.mark.parametrize('value', [0.0001, 1e-9, 2**-11]) +def test_scale_1024_small_number(value: float) -> None: + """Never scale a number below one up to a negative power.""" + assert converters.scale_1024(value, 9) == (value, 0) + + +@pytest.mark.parametrize('n_prefixes', [0, -1]) +def test_scale_1024_without_prefixes(n_prefixes: int) -> None: + """Never return a negative power when there are no prefixes.""" + assert converters.scale_1024(2048, n_prefixes) == (2048.0, 0) diff --git a/_python_utils_tests/test_formatters.py b/_python_utils_tests/test_formatters.py new file mode 100644 index 0000000..e683c88 --- /dev/null +++ b/_python_utils_tests/test_formatters.py @@ -0,0 +1,80 @@ +"""Tests for the formatting helpers in ``python_utils.formatters``.""" + +import datetime +import typing + +import pytest + +from python_utils import formatters + + +@pytest.mark.parametrize( + ('days', 'expected'), + [ + (30, '1 month ago'), + (35, '1 month and 5 days ago'), + (60, '2 months ago'), + (365, '1 year ago'), + (730, '2 years ago'), + ], +) +def test_timesince_units_do_not_overlap(days: int, expected: str) -> None: + """Take the weeks and days from what the larger units left over.""" + delta: datetime.timedelta = datetime.timedelta(days=days) + assert formatters.timesince(delta) == expected + + +@pytest.mark.parametrize( + ('delta', 'expected'), + [ + (datetime.timedelta(seconds=-1), '1 second ago'), + (datetime.timedelta(seconds=-61), '1 minute and 1 second ago'), + (datetime.timedelta(days=-400), '1 year and 1 month ago'), + ], +) +def test_timesince_negative_timedelta( + delta: datetime.timedelta, expected: str +) -> None: + """Describe a negative timedelta by its size, like a datetime.""" + assert formatters.timesince(delta) == expected + + +@pytest.mark.parametrize( + 'timezone', + [ + datetime.timezone.utc, + datetime.timezone(datetime.timedelta(hours=5, minutes=30)), + ], +) +def test_timesince_aware_datetime(timezone: datetime.timezone) -> None: + """Compare a timezone-aware datetime with an aware current time.""" + # Only the two largest units are shown, so the time this test takes + # cannot change the result. + age: datetime.timedelta = datetime.timedelta(hours=1, minutes=30) + moment: datetime.datetime = datetime.datetime.now(timezone) - age + assert formatters.timesince(moment) == '1 hour and 30 minutes ago' + + +@pytest.mark.parametrize( + ('name', 'expected'), + [ + ('SPAM_EGGS', 'spam_eggs'), + ('HTTP_OK', 'http_ok'), + ('MAX_SIZE_LIMIT', 'max_size_limit'), + ('HTML5Parser', 'html5_parser'), + ('toUTF8String', 'to_utf8_string'), + ('ABCD_e', 'abcd_e'), + ], +) +def test_camel_to_underscore_keeps_acronyms_whole( + name: str, expected: str +) -> None: + """Leave an acronym whole when an underscore or a digit follows it.""" + assert formatters.camel_to_underscore(name) == expected + + +def test_timesince_date_raises_type_error() -> None: + """Keep raising ``TypeError`` for a date, which has no time.""" + today: typing.Any = datetime.date.today() + with pytest.raises(TypeError): + formatters.timesince(today) diff --git a/_python_utils_tests/test_import.py b/_python_utils_tests/test_import.py index e78e117..bcdf773 100644 --- a/_python_utils_tests/test_import.py +++ b/_python_utils_tests/test_import.py @@ -1,7 +1,44 @@ """Tests for the import helpers in ``python_utils.import_``.""" +import collections.abc +import pathlib +import sys + +import pytest + from python_utils import import_, types +#: Top-level name of the package that the ``nested_module`` fixture creates. +NESTED_PACKAGE: str = 'python_utils_spam' + + +@pytest.fixture +def nested_module( + tmp_path: pathlib.Path, monkeypatch: pytest.MonkeyPatch +) -> collections.abc.Iterator[str]: + """Create ``python_utils_spam.eggs.bacon`` and yield its dotted name.""" + package: pathlib.Path = tmp_path / NESTED_PACKAGE + subpackage: pathlib.Path = package / 'eggs' + subpackage.mkdir(parents=True) + (package / '__init__.py').write_text('', encoding='utf-8') + (subpackage / '__init__.py').write_text('', encoding='utf-8') + (subpackage / 'bacon.py').write_text("ham = 'ham'\n", encoding='utf-8') + (subpackage / 'needy.py').write_text( + 'import python_utils_missing_dependency\n', encoding='utf-8' + ) + (subpackage / 'broken.py').write_text( + "raise AttributeError('broken on purpose')\n", encoding='utf-8' + ) + monkeypatch.setattr(sys, 'path', [str(tmp_path), *sys.path]) + + yield f'{NESTED_PACKAGE}.eggs.bacon' + + imported: list[str] = [ + name for name in sys.modules if name.split('.')[0] == NESTED_PACKAGE + ] + for name in imported: + del sys.modules[name] + def test_import_globals_relative_import() -> None: """Resolve relative imports across several levels.""" @@ -60,3 +97,109 @@ def test_import_locals_missing_module() -> None: 'python_utils.spam', exceptions=ImportError, globals_=globals() ) assert 'camel_to_underscore' in globals() + + +def test_import_global_nested_module(nested_module: str) -> None: + """Import a dotted name whose last module nothing imported before.""" + locals_: types.Dict[str, types.Any] = {} + globals_: types.Dict[str, types.Any] = {'__name__': __name__} + assert nested_module not in sys.modules + import_.import_global(nested_module, locals_=locals_, globals_=globals_) + assert globals_['ham'] == 'ham' + + +def test_import_global_missing_nested_module(nested_module: str) -> None: + """Keep the ``ImportError`` for a nested module that does not exist.""" + locals_: types.Dict[str, types.Any] = {} + globals_: types.Dict[str, types.Any] = {'__name__': __name__} + error: types.Any = import_.import_global( + f'{nested_module}.sausage', + exceptions=ImportError, + locals_=locals_, + globals_=globals_, + ) + assert type(error) is ImportError + assert str(error) == f'No module named {nested_module}.sausage' + + +def test_import_global_attribute_of_non_module() -> None: + """Report a path through an object that is not a module as missing.""" + locals_: types.Dict[str, types.Any] = {} + globals_: types.Dict[str, types.Any] = {'__name__': __name__} + error: types.Any = import_.import_global( + 'os.environ.spam', + exceptions=ImportError, + locals_=locals_, + globals_=globals_, + ) + assert type(error) is ImportError + assert str(error) == 'No module named os.environ.spam' + + +def test_import_global_nested_module_missing_dependency( + nested_module: str, +) -> None: + """Show the missing dependency of a nested module that does exist.""" + locals_: types.Dict[str, types.Any] = {} + globals_: types.Dict[str, types.Any] = {'__name__': __name__} + with pytest.raises(ModuleNotFoundError) as raised: + import_.import_global( + f'{NESTED_PACKAGE}.eggs.needy', locals_=locals_, globals_=globals_ + ) + + assert raised.value.name == 'python_utils_missing_dependency' + + +def test_import_global_nested_module_import_error(nested_module: str) -> None: + """Show the error that a nested module raises while it is imported.""" + locals_: types.Dict[str, types.Any] = {} + globals_: types.Dict[str, types.Any] = {'__name__': __name__} + with pytest.raises(AttributeError, match='broken on purpose'): + import_.import_global( + f'{NESTED_PACKAGE}.eggs.broken', locals_=locals_, globals_=globals_ + ) + + +def test_import_global_looks_an_attribute_up_once( + nested_module: str, tmp_path: pathlib.Path +) -> None: + """Ask a module for an attribute once, also when it computes it.""" + lazy: pathlib.Path = tmp_path / NESTED_PACKAGE / 'eggs' / 'lazy.py' + lazy.write_text( + 'calls = []\n\n\n' + 'def __getattr__(name):\n' + " if name == 'computed':\n" + ' calls.append(name)\n' + ' return calls\n' + ' raise AttributeError(name)\n', + encoding='utf-8', + ) + locals_: types.Dict[str, types.Any] = {} + globals_: types.Dict[str, types.Any] = {'__name__': __name__} + import_.import_global( + f'{NESTED_PACKAGE}.eggs.lazy.computed', + locals_=locals_, + globals_=globals_, + ) + + assert sys.modules[f'{NESTED_PACKAGE}.eggs.lazy'].calls == ['computed'] + + +def test_import_global_missing_below_detached_module( + nested_module: str, tmp_path: pathlib.Path +) -> None: + """Name the requested module when a module has another name itself.""" + detached: pathlib.Path = tmp_path / NESTED_PACKAGE / 'eggs' / 'detached.py' + detached.write_text( + "import types\n\nspam = types.ModuleType('python_utils_detached')\n", + encoding='utf-8', + ) + name: str = f'{NESTED_PACKAGE}.eggs.detached.spam.sausage' + locals_: types.Dict[str, types.Any] = {} + globals_: types.Dict[str, types.Any] = {'__name__': __name__} + error: types.Any = import_.import_global( + name, exceptions=ImportError, locals_=locals_, globals_=globals_ + ) + + assert type(error) is ImportError + assert str(error) == f'No module named {name}' diff --git a/_python_utils_tests/test_lazy_imports.py b/_python_utils_tests/test_lazy_imports.py index da8143e..54ec64c 100644 --- a/_python_utils_tests/test_lazy_imports.py +++ b/_python_utils_tests/test_lazy_imports.py @@ -11,6 +11,7 @@ import pytest import python_utils +from _python_utils_tests import clock def _run_clean(code: str) -> subprocess.CompletedProcess[str]: @@ -94,10 +95,13 @@ def test_star_import_resolves_all_names() -> None: @pytest.mark.asyncio -async def test_aio_timeout_generator_default_iterable() -> None: +async def test_aio_timeout_generator_default_iterable( + fake_clock: clock.FakeClock, +) -> None: """Default the iterable to ``aio.acount`` when omitted.""" # With no iterable the generator defaults to ``aio.acount`` -- exercising # the lazy ``aio``/``asyncio`` import and the None-resolution branch. + # The fake clock stands still, so the timeout cannot end the loop early. count = 0 generator: collections.abc.AsyncGenerator[object, None] = ( python_utils.aio_timeout_generator(timeout=0.05, interval=0.0) diff --git a/_python_utils_tests/test_logged.py b/_python_utils_tests/test_logged.py new file mode 100644 index 0000000..f89a1ae --- /dev/null +++ b/_python_utils_tests/test_logged.py @@ -0,0 +1,224 @@ +"""Tests for the logging mixins in ``python_utils.logger``.""" + +import logging + +import pytest + +from python_utils import logger + + +def records_of( + caplog: pytest.LogCaptureFixture, name: str +) -> list[logging.LogRecord]: + """Return the captured records of the logger called ``name``.""" + return [record for record in caplog.records if record.name == name] + + +def test_exception_logs_traceback(caplog: pytest.LogCaptureFixture) -> None: + """Attach the active exception to the record, like ``logging`` does.""" + + class Spam(logger.Logged): + pass + + spam: Spam = Spam() + name: str = Spam.logger.name + with caplog.at_level(logging.DEBUG, logger=name): + try: + int('eggs') + except ValueError: + spam.exception('bacon') + + record: logging.LogRecord = records_of(caplog, name)[-1] + assert record.exc_info is not None + assert record.exc_info[0] is ValueError + + +def test_exception_without_traceback(caplog: pytest.LogCaptureFixture) -> None: + """Leave the traceback out when ``exc_info`` is switched off.""" + + class Spam(logger.Logged): + pass + + spam: Spam = Spam() + name: str = Spam.logger.name + with caplog.at_level(logging.DEBUG, logger=name): + try: + int('eggs') + except ValueError: + spam.exception('bacon', exc_info=False) + + record: logging.LogRecord = records_of(caplog, name)[-1] + assert not record.exc_info + + +def test_records_name_the_caller(caplog: pytest.LogCaptureFixture) -> None: + """Report the calling function in every record, not the mixin.""" + + class Spam(logger.Logged): + pass + + spam: Spam = Spam() + name: str = Spam.logger.name + with caplog.at_level(logging.DEBUG, logger=name): + spam.debug('debug') + spam.info('info') + spam.warning('warning') + spam.error('error') + spam.critical('critical') + spam.exception('exception') + spam.log(logging.INFO, 'log') + + records: list[logging.LogRecord] = records_of(caplog, name) + assert len(records) == 7 + assert {record.funcName for record in records} == { + 'test_records_name_the_caller' + } + assert {record.filename for record in records} == {'test_logged.py'} + + +def log_one_frame_up(spam: logger.Logged) -> None: + """Log a record that points at the caller of this helper.""" + spam.info('eggs', stacklevel=2) + + +def test_stacklevel_counts_from_the_caller( + caplog: pytest.LogCaptureFixture, +) -> None: + """Count an explicit ``stacklevel`` from the calling function.""" + + class Spam(logger.Logged): + pass + + spam: Spam = Spam() + name: str = Spam.logger.name + with caplog.at_level(logging.DEBUG, logger=name): + log_one_frame_up(spam) + + record: logging.LogRecord = records_of(caplog, name)[-1] + assert record.funcName == 'test_stacklevel_counts_from_the_caller' + + +class Bacon: + """A base class whose ``__new__`` needs the constructor arguments.""" + + value: int + + def __new__(cls, value: int) -> 'Bacon': + """Store ``value`` on the new instance.""" + self: Bacon = super().__new__(cls) + self.value = value + return self + + +def test_new_forwards_arguments_to_builtin() -> None: + """Keep the value of a built-in type that is created in ``__new__``.""" + + class LoggedInt(logger.Logged, int): + pass + + class LoggedStr(logger.Logged, str): + pass + + number: LoggedInt = LoggedInt(5) + octal: LoggedInt = LoggedInt('7', base=8) + text: LoggedStr = LoggedStr('spam') + assert number == 5 + assert octal == 7 + assert text == 'spam' + + +def test_new_forwards_arguments_to_base() -> None: + """Pass the constructor arguments on to the next ``__new__``.""" + + class Eggs(logger.Logged, Bacon): + pass + + eggs: Eggs = Eggs(5) + assert eggs.value == 5 + assert Eggs.logger.name.endswith('.Eggs') + + +def test_new_accepts_arguments_for_init() -> None: + """Keep accepting arguments that only ``__init__`` uses.""" + + class Spam(logger.Logged): + def __init__(self, value: int, name: str = 'spam') -> None: + """Store both arguments on the instance.""" + self.value: int = value + self.name: str = name + + spam: Spam = Spam(1, name='eggs') + assert (spam.value, spam.name) == (1, 'eggs') + + +class Singleton: + """A base class whose ``__new__`` takes no constructor arguments.""" + + instance: 'Singleton | None' = None + + def __new__(cls) -> 'Singleton': + """Create the one instance on the first call and reuse it after.""" + if cls.instance is None: + cls.instance = super().__new__(cls) + + return cls.instance + + +def test_new_falls_back_without_arguments() -> None: + """Keep working with a base ``__new__`` that takes no arguments.""" + + class Config(logger.Logged, Singleton): + def __init__(self, path: str = 'defaults.ini') -> None: + """Store the path on the instance.""" + self.path: str = path + + config: Config = Config('settings.ini') + assert config.path == 'settings.ini' + assert Config('other.ini') is config + + +def test_new_reports_the_first_error() -> None: + """Raise the error of the forwarded call when no call works.""" + + class Eggs(logger.Logged, Bacon): + def __init__(self, *args: int) -> None: + """Accept any number of arguments.""" + + # The message is the interpreter's own and differs on PyPy. + with pytest.raises(TypeError): + Eggs(1, 2) + + +def test_logger_is_created_at_first_instance() -> None: + """Create the logger of a class when it is first instantiated. + + ``logging.config.dictConfig`` disables every logger that exists when it + runs. A logger that is created when the class is defined would be gone + for an application that configures logging after its imports. + """ + + class LateSpam(logger.Logged): + pass + + name: str = f'{LateSpam.__module__}.LateSpam' + assert name not in logging.Logger.manager.loggerDict + assert 'logger' not in vars(LateSpam) + + LateSpam() + assert name in logging.Logger.manager.loggerDict + assert LateSpam.logger.name == name + + +def test_subclass_inherits_class_body_logger() -> None: + """Let a subclass use the logger from the class body of its parent.""" + custom: logging.Logger = logging.getLogger('python_utils.tests.custom') + + class Parent(logger.Logged): + logger = custom + + class Child(Parent): + pass + + assert Child.logger is custom + Child() + assert Child.logger.name.endswith('.Child') diff --git a/_python_utils_tests/test_logger.py b/_python_utils_tests/test_logger.py index 15122b5..52eda47 100644 --- a/_python_utils_tests/test_logger.py +++ b/_python_utils_tests/test_logger.py @@ -1,12 +1,32 @@ # mypy: disable-error-code=misc """Tests for the loguru mixin in ``python_utils.loguru``.""" +import collections.abc + +import loguru as loguru_lib import pytest from python_utils import loguru pytest.importorskip('loguru') +#: Messages that ``str.format`` rejects or would rewrite. +BRACE_MESSAGES: tuple[str, ...] = ( + 'payload {"spam": 1}', + 'set {1, 2}', + 'unbalanced {', + 'positional {}', +) + + +@pytest.fixture +def messages() -> collections.abc.Iterator[list[str]]: + """Collect the text of every record loguru emits during a test.""" + collected: list[str] = [] + sink_id: int = loguru_lib.logger.add(collected.append, format='{message}') + yield collected + loguru_lib.logger.remove(sink_id) + def test_logurud() -> None: """Expose all loguru log-level methods on a subclass.""" @@ -22,3 +42,103 @@ class MyClass(loguru.Logurud): my_class.critical('critical') my_class.exception('exception') my_class.log(0, 'log') + + +@pytest.mark.parametrize('message', BRACE_MESSAGES) +def test_logurud_message_with_braces( + messages: list[str], message: str +) -> None: + """Log a message that contains braces exactly as it was given.""" + + class MyClass(loguru.Logurud): + pass + + my_class: loguru.Logurud = MyClass() + my_class.info(message) + assert messages == [f'{message}\n'] + + +def test_logurud_braces_in_every_method(messages: list[str]) -> None: + """Accept a message with braces in every log-level method.""" + + class MyClass(loguru.Logurud): + pass + + message: str = 'payload {"spam": 1}' + my_class: loguru.Logurud = MyClass() + my_class.debug(message) + my_class.info(message) + my_class.warning(message) + my_class.error(message) + my_class.critical(message) + my_class.exception(message) + my_class.log(20, message) + assert len(messages) == 7 + assert all(logged.startswith(message) for logged in messages) + + +def test_logurud_keeps_explicit_extra() -> None: + """Keep passing an explicit ``extra`` on to the loguru record.""" + + class MyClass(loguru.Logurud): + pass + + extras: list[dict[str, object]] = [] + sink_id: int = loguru_lib.logger.add( + lambda message: extras.append(dict(message.record['extra'])), + format='{message}', + ) + try: + my_class: loguru.Logurud = MyClass() + my_class.info('spam') + my_class.info('spam', extra={'eggs': 1}) + finally: + loguru_lib.logger.remove(sink_id) + + assert extras == [{}, {'extra': {'eggs': 1}}] + + +def test_logurud_new_forwards_arguments() -> None: + """Keep the value of a built-in type that is created in ``__new__``.""" + + class LoggedInt(loguru.Logurud, int): + pass + + number: loguru.Logurud = LoggedInt(5) + octal: loguru.Logurud = LoggedInt('7', base=8) + assert isinstance(number, int) + assert isinstance(octal, int) + assert number == 5 + assert octal == 7 + + +def test_logurud_new_falls_back_without_arguments() -> None: + """Keep working with a base ``__new__`` that takes no arguments.""" + + class Plain: + def __new__(cls) -> 'Plain': + """Create the instance without any constructor arguments.""" + return super().__new__(cls) + + class MyClass(loguru.Logurud, Plain): + def __init__(self, value: int) -> None: + """Store the value on the instance.""" + self.value: int = value + + my_class: loguru.Logurud = MyClass(5) + assert isinstance(my_class, MyClass) + assert my_class.value == 5 + + +def test_logurud_new_accepts_arguments_for_init() -> None: + """Keep accepting arguments that only ``__init__`` uses.""" + + class MyClass(loguru.Logurud): + def __init__(self, value: int, name: str = 'spam') -> None: + """Store both arguments on the instance.""" + self.value: int = value + self.name: str = name + + my_class: loguru.Logurud = MyClass(1, name='eggs') + assert isinstance(my_class, MyClass) + assert (my_class.value, my_class.name) == (1, 'eggs') diff --git a/_python_utils_tests/test_time.py b/_python_utils_tests/test_time.py index da26c9e..c9cca14 100644 --- a/_python_utils_tests/test_time.py +++ b/_python_utils_tests/test_time.py @@ -3,12 +3,18 @@ import asyncio import datetime import itertools +import time import pytest import python_utils +from _python_utils_tests import clock from python_utils import types +#: Far longer than every timeout in this module, so a generator that sleeps +#: this long is always interrupted first. +STALL: float = 10.0 + @pytest.mark.parametrize( 'timeout,interval,interval_multiplier,maximum_interval,iterable,result', @@ -28,6 +34,8 @@ ) @pytest.mark.asyncio async def test_aio_timeout_generator( + fake_clock: clock.FakeClock, + monkeypatch: pytest.MonkeyPatch, timeout: float, interval: float, interval_multiplier: float, @@ -36,6 +44,13 @@ async def test_aio_timeout_generator( result: int, ) -> None: """Stop the async generator near the configured timeout.""" + + async def sleep(delay: float) -> None: + """Let the fake clock pass the delay without waiting for it.""" + fake_clock.sleep(delay) + + monkeypatch.setattr(asyncio, 'sleep', sleep) + i = None async for i in python_utils.aio_timeout_generator( timeout, interval, iterable, maximum_interval=maximum_interval @@ -46,12 +61,13 @@ async def test_aio_timeout_generator( @pytest.mark.parametrize( - 'timeout,interval,interval_multiplier,maximum_interval,iterable,result', + 'timeout,interval,interval_multiplier,maximum_interval,iterable,result,' + 'sleeps', [ - (0.1, 0.06, 0.5, 0.1, 'abc', 'c'), - (0.1, 0.07, 0.5, 0.1, itertools.count, 2), - (0.1, 0.07, 0.5, 0.1, itertools.count(), 2), - (0.1, 0.06, 1.0, None, 'abc', 'c'), + (0.1, 0.06, 0.5, 0.1, 'abc', 'c', [0.06, 0.03, 0.015]), + (0.1, 0.07, 0.5, 0.1, itertools.count, 2, [0.07, 0.035]), + (0.1, 0.07, 0.5, 0.1, itertools.count(), 2, [0.07, 0.035]), + (0.1, 0.06, 1.0, None, 'abc', 'c', [0.06, 0.06]), ( datetime.timedelta(seconds=0.1), datetime.timedelta(seconds=0.06), @@ -59,10 +75,12 @@ async def test_aio_timeout_generator( datetime.timedelta(seconds=0.1), itertools.count, 2, + [0.06, 0.1], ), ], ) def test_timeout_generator( + fake_clock: clock.FakeClock, timeout: float, interval: float, interval_multiplier: float, @@ -73,8 +91,9 @@ def test_timeout_generator( types.Callable[..., types.Iterable[types.Any]], ], result: int, + sleeps: types.List[float], ) -> None: - """Stop the sync generator near the configured timeout.""" + """Stop the sync generator at the timeout and scale the interval.""" i = None for i in python_utils.timeout_generator( timeout=timeout, @@ -86,59 +105,97 @@ def test_timeout_generator( assert i is not None assert i == result + assert fake_clock.sleeps == pytest.approx(sleeps) -@pytest.mark.asyncio -async def test_aio_generator_timeout_detector() -> None: - """Raise or exit on per-item and total timeouts.""" - # Make pyright happy - i = None +def test_timeout_generator_real_clock() -> None: + """Keep yielding on the real clock until the timeout has passed.""" + timeout: float = 0.05 + interval: float = 0.01 + start: float = time.perf_counter() + items: types.List[int] = list( + python_utils.timeout_generator(timeout, interval, itertools.count()) + ) + elapsed: float = time.perf_counter() - start + + # A sleep can take longer than requested but never shorter, so these + # hold on any machine. The exact number of items does not. + assert items == list(range(len(items))) + assert len(items) <= timeout / interval + 2 + assert elapsed >= timeout + + +async def stalling_generator() -> types.AsyncGenerator[int, None]: + """Yield 0-4 without waiting, then stall before the next item.""" + for i in range(10): + if i == 5: + await asyncio.sleep(STALL) + yield i + + +def ticking_generator( + fake_clock: clock.FakeClock, +) -> types.AsyncGenerator[int, None]: + """Yield 0-9 and let 0.1 seconds pass on the fake clock for each item.""" async def generator() -> types.AsyncGenerator[int, None]: - """Yield 0-9 with increasing sleeps between items.""" + """Advance the fake clock before every item.""" for i in range(10): - await asyncio.sleep(i / 20.0) + fake_clock.sleep(0.1) yield i + return generator() + + +@pytest.mark.asyncio +async def test_aio_generator_timeout_detector( + fake_clock: clock.FakeClock, +) -> None: + """Raise or exit on per-item and total timeouts.""" + # Make pyright happy + i = None + detector = python_utils.aio_generator_timeout_detector # Test regular timeout with reraise with pytest.raises(asyncio.TimeoutError): - async for i in detector(generator(), 0.25): + async for i in detector(stalling_generator(), 0.05): pass # Test regular timeout with clean exit - async for i in detector(generator(), 0.25, on_timeout=None): + async for i in detector(stalling_generator(), 0.05, on_timeout=None): pass assert i == 4 # Test total timeout with reraise with pytest.raises(asyncio.TimeoutError): - async for i in detector(generator(), total_timeout=0.5): + async for i in detector( + ticking_generator(fake_clock), total_timeout=0.45 + ): pass # Test total timeout with clean exit - async for i in detector(generator(), total_timeout=0.5, on_timeout=None): + async for i in detector( + ticking_generator(fake_clock), total_timeout=0.45, on_timeout=None + ): pass assert i == 4 # Test stop iteration - async for i in detector(generator(), on_timeout=None): + async for i in detector(ticking_generator(fake_clock), on_timeout=None): pass + assert i == 9 + @pytest.mark.asyncio async def test_aio_generator_timeout_detector_decorator_reraise() -> None: """Reraise ``TimeoutError`` on a per-item timeout.""" - # Test regular timeout with reraise - @python_utils.aio_generator_timeout_detector_decorator(timeout=0.05) - async def generator_timeout() -> types.AsyncGenerator[int, None]: - """Yield with increasing delays to trip the timeout.""" - for i in range(10): - await asyncio.sleep(i / 100.0) - yield i + generator_timeout = python_utils.aio_generator_timeout_detector_decorator( + timeout=0.05 + )(stalling_generator) with pytest.raises(asyncio.TimeoutError): async for _ in generator_timeout(): @@ -152,14 +209,9 @@ async def test_aio_generator_timeout_detector_decorator_clean_exit() -> None: i = None # Test regular timeout with clean exit - @python_utils.aio_generator_timeout_detector_decorator( + generator_clean = python_utils.aio_generator_timeout_detector_decorator( timeout=0.05, on_timeout=None - ) - async def generator_clean() -> types.AsyncGenerator[int, None]: - """Yield with increasing delays to trip the timeout.""" - for i in range(10): - await asyncio.sleep(i / 100.0) - yield i + )(stalling_generator) async for i in generator_clean(): pass @@ -168,17 +220,16 @@ async def generator_clean() -> types.AsyncGenerator[int, None]: @pytest.mark.asyncio -async def test_aio_generator_timeout_detector_decorator_reraise_total() -> ( - None -): +async def test_aio_generator_timeout_detector_decorator_reraise_total( + fake_clock: clock.FakeClock, +) -> None: """Reraise ``TimeoutError`` on a total timeout.""" # Test total timeout with reraise - @python_utils.aio_generator_timeout_detector_decorator(total_timeout=0.1) + @python_utils.aio_generator_timeout_detector_decorator(total_timeout=0.45) async def generator_reraise() -> types.AsyncGenerator[int, None]: - """Yield with increasing delays to trip the timeout.""" - for i in range(10): - await asyncio.sleep(i / 100.0) + """Let the fake clock pass the total timeout while yielding.""" + async for i in ticking_generator(fake_clock): yield i with pytest.raises(asyncio.TimeoutError): @@ -187,19 +238,20 @@ async def generator_reraise() -> types.AsyncGenerator[int, None]: @pytest.mark.asyncio -async def test_aio_generator_timeout_detector_decorator_clean_total() -> None: +async def test_aio_generator_timeout_detector_decorator_clean_total( + fake_clock: clock.FakeClock, +) -> None: """Exit cleanly on total timeout when ``on_timeout`` is ``None``.""" # Make pyright happy i = None # Test total timeout with clean exit @python_utils.aio_generator_timeout_detector_decorator( - total_timeout=0.1, on_timeout=None + total_timeout=0.45, on_timeout=None ) async def generator_clean_total() -> types.AsyncGenerator[int, None]: - """Yield with increasing delays to trip the timeout.""" - for i in range(10): - await asyncio.sleep(i / 100.0) + """Let the fake clock pass the total timeout while yielding.""" + async for i in ticking_generator(fake_clock): yield i async for i in generator_clean_total(): diff --git a/conftest.py b/conftest.py new file mode 100644 index 0000000..4a22f5b --- /dev/null +++ b/conftest.py @@ -0,0 +1,8 @@ +""" +Load the shared fixtures for the tests and the doctests. + +The fixtures live in the tests package. They are loaded from the repository +root because the doctests in ``python_utils`` need them as well. +""" + +pytest_plugins: tuple[str, ...] = ('_python_utils_tests.clock',) diff --git a/pyproject.toml b/pyproject.toml index c875f23..ce32e4d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,7 +8,7 @@ module-root = '' module-name = 'python_utils' # Keep the tests and tox config in the sdist (parity with the old MANIFEST.in) # so downstream packagers can build and test from the source distribution. -source-include = ['_python_utils_tests/**/*.py', 'tox.ini'] +source-include = ['_python_utils_tests/**/*.py', 'conftest.py', 'tox.ini'] [project] name = 'python-utils' diff --git a/python_utils/converters.py b/python_utils/converters.py index ebe57ce..2206459 100644 --- a/python_utils/converters.py +++ b/python_utils/converters.py @@ -50,7 +50,13 @@ def to_int( When a (regexp) object (has a search method) is given, that will be used. WHen a string is given, re.compile will be run over it first - The last group of the regexp will be used as value + The last group of the regexp will be used as value. A regexp without a + group uses the whole match. + + With ``regexp=True`` only the first run of digits is used. A sign, a + decimal point and an exponent are not part of it, so the result can differ + from the same call without a regexp. A regexp can only search a ``str``. + Any other input gives the default. >>> to_int('abc') 0 @@ -94,6 +100,14 @@ def to_int( 1234 >>> to_int('abc', default=1) 1 + >>> to_int('abc123', regexp=r'\d+') + 123 + >>> to_int('-5'), to_int('-5', regexp=True) + (-5, 5) + >>> to_int('1e3'), to_int('1e3', regexp=True) + (0, 1) + >>> to_int(123), to_int(123, regexp=True) + (123, 0) >>> to_int('abc', regexp=123) Traceback (most recent call last): ... @@ -110,7 +124,10 @@ def to_int( try: if regexp and input_ and (match := regexp.search(input_)): - input_ = match.groups()[-1] + # A pattern without a capture group has no last group, so the + # whole match is the value. + groups: tuple[str | None, ...] = match.groups() + input_ = groups[-1] if groups else match.group(0) if input_ is None: return default @@ -127,10 +144,12 @@ def to_float( regexp: _RegexpType = None, ) -> _aliases.Number: r""" - Convert the given `input_` to an integer or return default. + Convert the given `input_` to a float or return default. When trying to convert the exceptions given in the exception parameter - are automatically caught and the default will be returned. + are automatically caught and the default will be returned. The default is + the int ``0``, so a failed conversion returns an ``int`` unless a float + is passed as default. The regexp parameter allows for a regular expression to find the digits in a string. @@ -138,8 +157,31 @@ def to_float( When a (regexp) object (has a search method) is given, that will be used. When a string is given, re.compile will be run over it first - The last group of the regexp will be used as value + The first group of the regexp will be used as value. A regexp without a + group uses the whole match. + With ``regexp=True`` only the first run of digits is used, with the + decimal part that follows it. A sign, a leading decimal point and an + exponent are not part of it, so the result can differ from the same call + without a regexp. A regexp can only search a ``str``. Any other input + gives the default. + + >>> to_float('abc') + 0 + >>> to_float('abc', default=0.0) + 0.0 + >>> to_float('abc1.5', regexp=r'\d+\.\d+') + 1.5 + >>> to_float('a1b2', regexp=r'(\d)\D(\d)') + 1.0 + >>> to_float('-1.5'), to_float('-1.5', regexp=True) + (-1.5, 1.5) + >>> to_float('.5'), to_float('.5', regexp=True) + (0.5, 5.0) + >>> to_float('1e3'), to_float('1e3', regexp=True) + (1000.0, 1.0) + >>> to_float(1.5), to_float(1.5, regexp=True) + (1.5, 0) >>> '%.2f' % to_float('abc') '0.00' >>> '%.2f' % to_float('1') @@ -188,7 +230,9 @@ def to_float( try: if regexp and (match := regexp.search(input_)): - input_ = match.group(1) + # A pattern without a capture group has no first group, so the + # whole match is the value. + input_ = match.group(1) if match.groups() else match.group(0) return float(input_) except exception: return default @@ -231,7 +275,7 @@ def to_str( ) -> bytes: """Convert objects to string, encodes to the given encoding. - :rtype: str + :rtype: bytes >>> to_str('a') b'a' @@ -242,9 +286,9 @@ def to_str( >>> class Foo(object): ... __str__ = lambda s: 'a' >>> to_str(Foo()) - 'a' + b'a' >>> to_str(Foo) - "" + b"" """ if not isinstance(input_, bytes): if not hasattr(input_, 'encode'): @@ -278,7 +322,9 @@ def scale_1024( if x <= 0: power = 0 else: - power = min(int(math.log(x, 2) / 10), n_prefixes - 1) + # Never below zero: a number under 1 has a negative logarithm, and a + # negative power would index the prefixes from the wrong end. + power = max(min(int(math.log(x, 2) / 10), n_prefixes - 1), 0) scaled = float(x) / (2 ** (10 * power)) return scaled, power @@ -386,11 +432,20 @@ def remap( # pyright: ignore[reportInconsistentOverload] If floating point remaps need to be done my suggestion is to pass at least one parameter as a `decimal.Decimal`. This will ensure that the output - from this function is accurate. I left passing `floats` for backwards - compatibility and there is no conversion done from float to - `decimal.Decimal` unless one of the passed parameters has a type of - `decimal.Decimal`. This will ensure that any existing code that uses this - function will work exactly how it has in the past. + from this function is accurate, as long as every value with a fraction is + a `decimal.Decimal` itself. A `float` is converted exactly, with its + binary rounding error included: + + >>> remap(0.1, 0, 1, 0, decimal.Decimal(10)) + Decimal('1.000000000000000055511151231') + >>> remap(decimal.Decimal('0.1'), 0, 1, 0, 10) + Decimal('1.0') + + I left passing `floats` for backwards compatibility and there is no + conversion done from float to `decimal.Decimal` unless one of the passed + parameters has a type of `decimal.Decimal`. This will ensure that any + existing code that uses this function will work exactly how it has in the + past. Some edge cases to test >>> remap(1, 0, 0, 1, 2) diff --git a/python_utils/formatters.py b/python_utils/formatters.py index 667b376..71082c4 100644 --- a/python_utils/formatters.py +++ b/python_utils/formatters.py @@ -42,9 +42,10 @@ def camel_to_underscore(name: str) -> str: # Uppercase and the previous character isn't upper/underscore? # Add the underscore output.append('_') - elif i > 3 and not c.isupper(): + elif i > 3 and c.islower(): # Will return the last 3 letters to check if we are changing - # case + # case. Only a lowercase letter ends an acronym, an underscore + # or a digit after it leaves the acronym whole. previous = name[i - 3 : i] if previous.isalpha() and previous.isupper(): output.insert(len(output) - 1, '_') @@ -140,16 +141,32 @@ def timesince( '1 hour and 2 minutes ago' """ if isinstance(dt, datetime.timedelta): - diff = dt + # A negative timedelta has negative days and positive seconds, so + # only its size can be described. + diff = abs(dt) else: - now = datetime.datetime.now() + # An aware datetime can only be compared with an aware current time. + # For a naive datetime `tzinfo` is `None`, which gives local time. + # A date has no `tzinfo`, and no time to subtract either. It gets + # the same `TypeError` from the subtraction as it always did. + now = datetime.datetime.now(getattr(dt, 'tzinfo', None)) diff = abs(now - dt) + # Every unit takes its share from what the larger units left over, so a + # day is never counted twice. + years: int + months: int + weeks: int + days: int + years, days = divmod(diff.days, 365) + months, days = divmod(days, 30) + weeks, days = divmod(days, 7) + periods = ( - (diff.days / 365, 'year', 'years'), - (diff.days % 365 / 30, 'month', 'months'), - (diff.days % 30 / 7, 'week', 'weeks'), - (diff.days % 7, 'day', 'days'), + (years, 'year', 'years'), + (months, 'month', 'months'), + (weeks, 'week', 'weeks'), + (days, 'day', 'days'), (diff.seconds / 3600, 'hour', 'hours'), (diff.seconds % 3600 / 60, 'minute', 'minutes'), (diff.seconds % 60, 'second', 'seconds'), diff --git a/python_utils/import_.py b/python_utils/import_.py index faed863..1b070ba 100644 --- a/python_utils/import_.py +++ b/python_utils/import_.py @@ -12,6 +12,8 @@ relative imports and custom exception handling. """ +import importlib +import types import typing from python_utils import _aliases @@ -25,6 +27,46 @@ class DummyError(Exception): DummyException = DummyError +def _get_attribute(module: typing.Any, attr: str, name: str) -> typing.Any: + """Return ``module.attr``, importing it as a submodule when needed. + + A submodule only becomes an attribute of its parent once something has + imported it. + + Args: + module: The module, or other object, to take the attribute from. + attr: The name of the attribute or submodule. + name: The full dotted name that is being imported, for the error. + + Returns: + The attribute or the imported submodule. + + Raises: + ImportError: When ``module`` has no such attribute or submodule. An + error that the submodule raises while it is imported is passed + on as it is, a missing dependency included. + """ + try: + return getattr(module, attr) + except AttributeError as error: + if not isinstance(module, types.ModuleType): + # The same error as for a missing module, as it always was. + raise ImportError( # noqa: TRY004 + f'No module named {name}' + ) from error + + submodule: str = f'{module.__name__}.{attr}' + try: + return importlib.import_module(submodule) + except ModuleNotFoundError as error: + missing: str = error.name or '' + if missing != submodule and not submodule.startswith(f'{missing}.'): + # The submodule exists and one of its own imports is missing. + raise + + raise ImportError(f'No module named {name}') from error + + def import_global( # noqa: C901 name: str, modules: list[str] | None = None, @@ -41,7 +83,11 @@ def import_global( # noqa: C901 Args: name (str): the name of the module to import, e.g. sys - modules (str): the modules to import, use None for everything + modules (list[str]): the names to import from the module, use None + for everything. An empty list also imports everything. A single + name needs a list as well, a bare string is read as a collection + of one-character names. Names that start with an underscore are + never imported. exceptions (Exception): the exception to catch, e.g. ImportError locals_: the `locals()` method (in case you need a different scope) globals_: the `globals()` method (in case you need a different scope) @@ -82,13 +128,8 @@ def import_global( # noqa: C901 # Make sure we get the right part of a dotted import (i.e. # spam.eggs should return eggs, not spam) - try: - for attr in name_parts[1:]: - module = getattr(module, attr) - except AttributeError as e: - raise ImportError( - 'No module named ' + '.'.join(name_parts) - ) from e + for attr in name_parts[1:]: + module = _get_attribute(module, attr, '.'.join(name_parts)) # If no list of modules is given, autodetect from either __all__ # or a dir() of the module diff --git a/python_utils/logger.py b/python_utils/logger.py index f231290..2fb54e3 100644 --- a/python_utils/logger.py +++ b/python_utils/logger.py @@ -30,6 +30,7 @@ import abc import collections.abc import logging +import sys import types import typing @@ -55,6 +56,10 @@ | BaseException | None ) +#: Extra frames between ``Logger.exception`` and its caller. Up to Python 3.10 +#: ``Logger.exception`` calls ``Logger.error``, and ``logging`` counts that +#: frame against the ``stacklevel``. From Python 3.11 it skips its own frames. +_EXCEPTION_FRAMES: int = int(sys.version_info < (3, 11)) #: Parameter specification capturing a wrapped logger method's arguments. _P = typing.ParamSpec('_P') #: Covariant return-type variable for wrapped logger methods. @@ -148,6 +153,45 @@ def log( """Log ``msg`` at the integer ``level``.""" +def _create_instance( + new: collections.abc.Callable[..., _T], + cls: type[_T], + args: tuple[typing.Any, ...], + kwargs: dict[str, typing.Any], +) -> _T: + """Create an instance with the next ``__new__`` in line. + + The constructor arguments are passed on, so that a class which is created + in ``__new__`` gets its value, as ``int`` and ``str`` do. A ``__new__`` + that does not take them is called without, the way it always was. That + second call repeats whatever the first one did before it raised. + + Args: + new: The next ``__new__`` in the method resolution order. + cls: The class to create an instance of. + args: The positional constructor arguments. + kwargs: The keyword constructor arguments. + + Returns: + The new instance. + + Raises: + TypeError: When ``new`` accepts the arguments in neither form. The + error is the one for the call with the arguments. + """ + if new is object.__new__: + # `object.__new__` takes no arguments, they are for `__init__`. + return new(cls) + + try: + return new(cls, *args, **kwargs) + except TypeError as error: + try: + return new(cls) + except TypeError: + raise error from None + + class LoggerBase(abc.ABC): """Class which automatically adds logging utilities to your class when inheriting. Expects `logger` to be a logging.Logger or compatible instance. @@ -182,6 +226,37 @@ def __get_name( # pyright: ignore[reportUnusedFunction] """Join the non-empty, stripped ``name_parts`` into a dotted name.""" return '.'.join(n.strip() for n in name_parts if n.strip()) + @classmethod + def _log_kwargs( + cls, + exc_info: _ExcInfoType, + stack_info: bool, + stacklevel: int, + extra: collections.abc.Mapping[str, object] | None, + ) -> dict[str, typing.Any]: + """Build the keyword arguments that every log method forwards. + + Subclasses with a logger that does not take the ``logging`` keyword + arguments can override this method. + + Args: + exc_info: Exception information to attach to the record. + stack_info: Whether to attach the current stack to the record. + stacklevel: How many frames up the caller of the log method is. + extra: Extra attributes for the record. + + Returns: + The keyword arguments for the method of ``cls.logger``. + """ + return { + 'exc_info': exc_info, + 'stack_info': stack_info, + # One level more, so the record names the caller of the log + # method instead of the log method itself. + 'stacklevel': stacklevel + 1, + 'extra': extra, + } + @decorators.wraps_classmethod(logging.Logger.debug) @classmethod def debug( @@ -197,10 +272,7 @@ def debug( return cls.logger.debug( # type: ignore[no-any-return] msg, *args, - exc_info=exc_info, - stack_info=stack_info, - stacklevel=stacklevel, - extra=extra, + **cls._log_kwargs(exc_info, stack_info, stacklevel, extra), ) @decorators.wraps_classmethod(logging.Logger.info) @@ -218,10 +290,7 @@ def info( return cls.logger.info( # type: ignore[no-any-return] msg, *args, - exc_info=exc_info, - stack_info=stack_info, - stacklevel=stacklevel, - extra=extra, + **cls._log_kwargs(exc_info, stack_info, stacklevel, extra), ) @decorators.wraps_classmethod(logging.Logger.warning) @@ -239,10 +308,7 @@ def warning( return cls.logger.warning( # type: ignore[no-any-return] msg, *args, - exc_info=exc_info, - stack_info=stack_info, - stacklevel=stacklevel, - extra=extra, + **cls._log_kwargs(exc_info, stack_info, stacklevel, extra), ) @decorators.wraps_classmethod(logging.Logger.error) @@ -260,10 +326,7 @@ def error( return cls.logger.error( # type: ignore[no-any-return] msg, *args, - exc_info=exc_info, - stack_info=stack_info, - stacklevel=stacklevel, - extra=extra, + **cls._log_kwargs(exc_info, stack_info, stacklevel, extra), ) @decorators.wraps_classmethod(logging.Logger.critical) @@ -281,10 +344,7 @@ def critical( return cls.logger.critical( # type: ignore[no-any-return] msg, *args, - exc_info=exc_info, - stack_info=stack_info, - stacklevel=stacklevel, - extra=extra, + **cls._log_kwargs(exc_info, stack_info, stacklevel, extra), ) @decorators.wraps_classmethod(logging.Logger.exception) @@ -293,7 +353,7 @@ def exception( cls, msg: object, *args: object, - exc_info: _ExcInfoType = None, + exc_info: _ExcInfoType = True, stack_info: bool = False, stacklevel: int = 1, extra: collections.abc.Mapping[str, object] | None = None, @@ -302,10 +362,12 @@ def exception( return cls.logger.exception( # type: ignore[no-any-return] msg, *args, - exc_info=exc_info, - stack_info=stack_info, - stacklevel=stacklevel, - extra=extra, + **cls._log_kwargs( + exc_info, + stack_info, + stacklevel + _EXCEPTION_FRAMES, + extra, + ), ) @decorators.wraps_classmethod(logging.Logger.log) @@ -325,10 +387,7 @@ def log( level, msg, *args, - exc_info=exc_info, - stack_info=stack_info, - stacklevel=stacklevel, - extra=extra, + **cls._log_kwargs(exc_info, stack_info, stacklevel, extra), ) @@ -382,4 +441,4 @@ def __new__( cls.logger = logging.getLogger( cls.__get_name(cls.__module__, cls.__name__) ) - return super().__new__(cls) + return _create_instance(super().__new__, cls, args, kwargs) diff --git a/python_utils/loguru.py b/python_utils/loguru.py index 31b838b..71b06c6 100644 --- a/python_utils/loguru.py +++ b/python_utils/loguru.py @@ -16,6 +16,7 @@ from __future__ import annotations +import collections.abc import typing import loguru @@ -35,6 +36,36 @@ class Logurud(logger_module.LoggerBase): logger: loguru.Logger + @classmethod + def _log_kwargs( + cls, + exc_info: object, + stack_info: bool, + stacklevel: int, + extra: collections.abc.Mapping[str, object] | None, + ) -> dict[str, typing.Any]: + """Forward only an explicit ``extra`` to the `loguru` logger. + + `loguru` runs ``str.format`` over the message as soon as a call has + any arguments. The `logging` keyword arguments would turn every + message with a brace in it into a broken format string, so they are + left out. An ``extra`` that the caller passed is still forwarded, as + it always reached the `loguru` record. + + Args: + exc_info: Ignored, `loguru` does not take this argument. + stack_info: Ignored, `loguru` does not take this argument. + stacklevel: Ignored, `loguru` does not take this argument. + extra: Forwarded when it is not ``None``. + + Returns: + The keyword arguments for the `loguru` logger. + """ + if extra is None: + return {} + + return {'extra': extra} + def __new__(cls, *args: typing.Any, **kwargs: typing.Any) -> Logurud: """ Creates a new instance of `Logurud` and initializes the `loguru` @@ -50,4 +81,6 @@ def __new__(cls, *args: typing.Any, **kwargs: typing.Any) -> Logurud: # `logger` is already declared at class scope; assign without # re-annotating to avoid an obscured-declaration error. cls.logger = loguru.logger.opt(depth=1) - return super().__new__(cls) + return logger_module._create_instance( # pyright: ignore[reportPrivateUsage] + super().__new__, cls, args, kwargs + )