How to patch a function of a global singleton to return a specified value given a specified input?

Viewed 104

File c.py has a global variable COBJ of an instance of C (Singleton)

class Singleton(type):
    _instances = {}
    def __call__(cls, *args, **kwargs):
        if cls not in cls._instances:
            cls._instances[cls] = super(Singleton, cls).__call__(*args, **kwargs)
        return cls._instances[cls]

class C(metaclass=Singleton):
    def __init__(self, a, b):
        #...
    def get(self, x):  # To be mocked/patched
        return ....

COBJ = C(1, 'z')

And I have a file x.py which import c.py and need to be tested.

from c import COBJ

class X:  # to be tested
    def f(self, a):
        x = COBJ.get(a)
        return x + '-AddedInX.f'  # just a simple example here

How to mock COBJ.get(a) for some specified parameter inputs.

import x

def test_f:
    xobj = X()
    input = 'abc' 
    # need to patch COBJ.get() to return '###' given 'abc'
    ...
    result = xobj.f(input)
    assert result == '###-AddedInX.f'

How to patch COJB.get() return a specified value given a specified input? Is COBJ pythonic way to create a singleton global object?

1 Answers

The following program will satisfy adding a patch to a function. When the function executes, it will return the corresponding value according to the value of the parameter, and a function can add multiple patches, and also add a method to delete the patch.

import inspect
from functools import wraps

class Patch:

    class _Empty:
        pass

    # Used to store the corresponding value.
    # Basic format: {obj_id: {func_name: {inp_v: ret_v}}}
    _patch_mapping = {}
    @classmethod
    def add_patch(cls, obj, f, inp_v, ret_v):
        """
        obj: The instance of the class to which the function belongs
        f: The function object that adds the patch
        inp_v: The parameter value entered when the function is executed
        ret_v: When the function receives a parameter value equal to `inp_v`, it will return `ret_v`
        """
        obj_id = id(obj)
        f = cls._wrapper(obj_id, f)
        cls._patch_mapping.setdefault(
                obj_id, {}
            ).setdefault(f.__name__, {})[inp_v] = ret_v
        return f
    
    @classmethod
    def remove_patch(cls, obj, f, inp_v=None):
        """Delete the patch, when `inp_v` is None, clear all the patches corresponding to the function."""
        obj_m = cls._patch_mapping.get(id(obj), {})
        if inp_v is None:
            obj_m.pop(f.__name__, None)
        else:
            obj_m.get(f.__name__, {}).pop(inp_v, None)

    @classmethod
    def _check(cls, func):
        """Check if the number of arguments to a function is 1 and it is a positional argument."""
        sig = inspect.signature(func)
        assert len(sig.parameters) == 1
        for j in sig.parameters.values():
            assert j.kind in (
                    inspect.Parameter.POSITIONAL_OR_KEYWORD, 
                    inspect.Parameter.POSITIONAL_ONLY
                )   

    @classmethod
    def _wrapper(cls, obj_id, func):
        cls._check(func)
        func_name = func.__name__
        @wraps(func)
        def inner(param):
            nonlocal obj_id, func_name
            """
            Judging whether it exists, if it exists, it returns the corresponding value, 
            and if it does not exist, the function is executed.
            """
            if (v := cls._patch_mapping.get(obj_id, {}).get(
                    func_name, {}
                ).get(param, cls._Empty)
            ) != cls._Empty:
                return v
            return func(param)
        return inner  


class Singleton(type):
    _instances = {}

    def add_patch(self, f, inp_v, ret_v):
        self.__dict__[f.__name__] = Patch.add_patch(self, f, inp_v, ret_v)

    def remove_patch(self, f, inp_v=None):
        Patch.remove_patch(self, f, inp_v)

    def __new__(cls, name, bases, attrs):
        # Add method to class.
        attrs["add_patch"] = cls.add_patch
        attrs["remove_patch"] = cls.remove_patch
        return super().__new__(cls, name, bases, attrs)

    def __call__(cls, *args, **kwargs):
        if cls not in cls._instances:
            cls._instances[cls] = super(Singleton, cls).__call__(*args, **kwargs)
        return cls._instances[cls]


class C(metaclass=Singleton):
    def __init__(self, a, b):
        #...
        pass

    def get(self, x):  # To be mocked/patched
        return "get"

COBJ = C(1, 'z')
print(COBJ.get("a"))
COBJ.add_patch(COBJ.get, "a", 2)
print(COBJ.get("a"))

COBJ1 = C(1, 'z')
print(COBJ1.get("a"))
COBJ1.add_patch(COBJ1.get, "a", 7)
COBJ1.add_patch(COBJ1.get, "b", 8)
print(COBJ1.get("a"))
print(COBJ1.get("b"))
COBJ1.remove_patch(COBJ1.get)
print(COBJ1.get("a"))
print(COBJ1.get("b"))

Output:

get
2
2
7
8
get
get
Related