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