Skip to content
Merged
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
30 changes: 15 additions & 15 deletions returns/contrib/pytest/plugin.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,19 +180,19 @@ def _patched_error_handler(
) -> _FunctionType:
if inspect.iscoroutinefunction(original):

async def wrapper(self, *args, **kwargs):
async def async_wrapper(self, *args, **kwargs):
original_result = await original(self, *args, **kwargs)
errs[id(original_result)] = original_result
return original_result

else:
return wraps(original)(async_wrapper) # type: ignore

def wrapper(self, *args, **kwargs):
original_result = original(self, *args, **kwargs)
errs[id(original_result)] = original_result
return original_result
def sync_wrapper(self, *args, **kwargs):
original_result = original(self, *args, **kwargs)
errs[id(original_result)] = original_result
return original_result

return wraps(original)(wrapper) # type: ignore
return wraps(original)(sync_wrapper) # type: ignore


def _patched_error_copier(
Expand All @@ -201,21 +201,21 @@ def _patched_error_copier(
) -> _FunctionType:
if inspect.iscoroutinefunction(original):

async def wrapper(self, *args, **kwargs):
async def async_wrapper(self, *args, **kwargs):
original_result = await original(self, *args, **kwargs)
if id(self) in errs:
errs[id(original_result)] = original_result
return original_result

else:
return wraps(original)(async_wrapper) # type: ignore

def wrapper(self, *args, **kwargs):
original_result = original(self, *args, **kwargs)
if id(self) in errs:
errs[id(original_result)] = original_result
return original_result
def sync_wrapper(self, *args, **kwargs):
original_result = original(self, *args, **kwargs)
if id(self) in errs:
errs[id(original_result)] = original_result
return original_result

return wraps(original)(wrapper) # type: ignore
return wraps(original)(sync_wrapper) # type: ignore


_ERROR_HANDLING_PATCHERS: Final = MappingProxyType({
Expand Down