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
16 changes: 15 additions & 1 deletion src/testfixtures/replace.py
Original file line number Diff line number Diff line change
Expand Up @@ -189,8 +189,22 @@ def on_class(self, attribute: Callable, replacement: Any, name: str | None = Non
if not callable(attribute):
name_text = f' named {name!r} ' if name else ' '
raise TypeError(f'attribute{name_text}must be a method')

container = None
if isinstance(attribute, classmethod_type):

qualname = getattr(attribute, '__qualname__', None)
if qualname and '<' not in qualname:
module_name = getattr(attribute, '__module__', None)
if module_name:
resolved = resolve(f"{module_name}.{qualname}")
if not resolved.found is not_there:
container = resolved.container
if resolved.found is not attribute:
container = None

if container is not None:
pass
elif isinstance(attribute, classmethod_type):
for referred in get_referents(attribute):
if isinstance(referred, class_type):
container = referred
Expand Down
42 changes: 42 additions & 0 deletions tests/sample1.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@
"""

from datetime import datetime, date
from functools import wraps
from typing import Any


def str_now_1():
Expand Down Expand Up @@ -77,3 +79,43 @@ class Slotted:
def __init__(self, x, y):
self.x = x
self.y = y


def wrap(fn: Any) -> Any:
@wraps(fn)
def inner(self: Any) -> Any:
return fn(self)
return inner


class ClassWithWrappedMethod:
@wrap
def method(self) -> str:
return 'original wrapped'


class Outer:
class Inner:
def method(self) -> str:
return 'original inner'

@staticmethod
def static() -> str:
return 'original static'

@classmethod
def class_(cls) -> str:
return 'original class'


def _instantiate(cls: type) -> Any:
return cls()


# Module-level name Singleton bound to the instance:
@_instantiate
class Singleton:
__slots__ = ()

def method(self) -> str:
return 'original singleton'
19 changes: 19 additions & 0 deletions tests/temporary_module.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,19 @@
import importlib
import sys
from contextlib import contextmanager
from types import ModuleType
from typing import Iterator

from testfixtures import TempDir, Replace


@contextmanager
def temporary_module(name: str, content: str) -> Iterator[ModuleType]:
with TempDir() as d:
d.write(f'{name}.py', content)
with Replace(target=sys.path, container=sys, name='path', replacement=[d.as_string()]):
try:
yield importlib.import_module(name)
finally:
if name in sys.modules:
del sys.modules[name]
122 changes: 115 additions & 7 deletions tests/test_replace.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,12 @@
import importlib
import inspect
import os
import sys
from contextlib import contextmanager
from gc import get_referrers
from operator import getitem
from typing import Any, Callable, Iterator
from unittest import TestCase

from testfixtures import (
Replacer,
Expand All @@ -13,18 +21,17 @@
replace_in_module,
ShouldWarn,
)
from unittest import TestCase

import os

from testfixtures.compat import PY_313_PLUS
from testfixtures.mock import Mock
from tests import sample1, sample3
from tests import sample2
from .sample1 import z, X
from .sample1 import z, X, Outer, ClassWithWrappedMethod, Singleton
from .sample3 import SOME_CONSTANT
from testfixtures.compat import PY_313_PLUS
from .temporary_module import temporary_module

from warnings import catch_warnings
# can't import the replace module simply as replace-the-function is imported into the main
# testfixtures package :-(
replace_module = importlib.import_module('testfixtures.replace')


class TestReplace(TestCase):
Expand Down Expand Up @@ -1423,3 +1430,104 @@ def test_mock_and_name(self):
# their attributes are not in it, confusing Replace, which deletes
# the attribute on restore as a result:
assert my_obj.foo is foo


class TestReplaceOnClassUsesQualname:
@staticmethod
@contextmanager
def check(
expression: Callable[[], Any], expected_get_referrers_calls: int = 0
) -> Iterator[Replacer]:
mock_get_referrers = Mock(side_effect=get_referrers)
with replace_in_module(get_referrers, mock_get_referrers, replace_module):
before = expression()
with Replacer() as replace:
yield replace
after = expression()
compare(expected=before, actual=after)
compare(mock_get_referrers.call_count, expected=expected_get_referrers_calls)

def test_simple_method(self) -> None:
expression = lambda: X().y()
with self.check(expression) as replace_:
replace_.on_class(X.y, lambda self_: 'mock y')
compare(expression(), expected='mock y')

def test_class_method(self) -> None:
expression = lambda: X().aMethod()
with self.check(expression) as replace_:
replace_.on_class(X.aMethod, lambda cls: 'mock aMethod')
compare(expression(), expected='mock aMethod')

def test_static_method(self) -> None:
expression = lambda: X().bMethod()
with self.check(expression) as replace_:
replace_.on_class(X.bMethod, lambda: 3)
compare(expression(), expected=3)

def test_inner_simple_method(self) -> None:
expression = lambda: Outer.Inner().method()
with self.check(expression) as replace_:
replace_.on_class(Outer.Inner.method, lambda self_: 'mock inner simple')
compare(expression(), expected='mock inner simple')

def test_inner_static_method(self) -> None:
expression = lambda: Outer.Inner().static()
with self.check(expression) as replace_:
replace_.on_class(Outer.Inner.static, lambda: 'mock inner static')
compare(expression(), expected='mock inner static')

def test_inner_class_method(self) -> None:
expression = lambda: Outer.Inner().class_()
with self.check(expression) as replace_:
replace_.on_class(Outer.Inner.class_, lambda cls: 'mock inner class')
compare(expression(), expected='mock inner class')

def test_wrapped(self) -> None:
expression = lambda: ClassWithWrappedMethod().method()
with self.check(expression) as replace_:
replace_.on_class(ClassWithWrappedMethod.method, lambda self: 'wrapped')
compare(expression(), expected='wrapped')

def test_wrapped_manually_unwrapped(self) -> None:
expression = lambda: ClassWithWrappedMethod().method()
with self.check(expression, expected_get_referrers_calls=1) as replace_:
t = ClassWithWrappedMethod.method.__wrapped__
with ShouldRaise(
AttributeError(f"could not find container of {repr(t)} using name 'method'")
):
replace_.on_class(t, lambda self: 'wrapped')
compare(expression(), expected='original wrapped')

def test_method_on_class_from_module_no_longer_in_sys_modules(self):
with temporary_module('unloaded', inspect.getsource(X)) as module:
expression = lambda: UnloadedX().y()

UnloadedX = module.X
del sys.modules['unloaded']

with self.check(expression, expected_get_referrers_calls=2) as replace_:
replace_.on_class(UnloadedX.y, lambda self: 'replaced')
compare(expression(), expected='replaced')
assert 'unloaded' in sys.modules, 'module was not reloaded'

def test_method_on_class_whose_name_is_bound_to_an_instance(self):
expression = lambda: Singleton.method()
with self.check(expression, expected_get_referrers_calls=2) as replace_:
replace_.on_class(type(Singleton).method, lambda self: 'replaced')
compare(expression(), expected='replaced')
compare(expression(), expected='original singleton')

def test_method_on_class_from_before_module_reload(self):
with temporary_module('reloaded', inspect.getsource(X)) as module:
expression = lambda: OldX().y()

OldX = module.X
importlib.reload(module)

with self.check(expression, expected_get_referrers_calls=2) as replace_:
replace_.on_class(OldX.y, lambda self: 'replaced')
compare(expression(), expected='replaced')
compare(module.X().y(), expected='original y')

compare(module.X().y(), expected='original y')
Loading