From c768c62ba173fb3e8a1b83d465cb67110551efa0 Mon Sep 17 00:00:00 2001 From: ShubhamDesai <42180509+ShubhamDesai@users.noreply.github.com> Date: Sat, 10 May 2025 13:55:30 -0400 Subject: [PATCH 1/2] Implement PEP 302 optional get_code loader method --- src/_pytest/assertion/rewrite.py | 24 +++++++++++++++--------- 1 file changed, 15 insertions(+), 9 deletions(-) diff --git a/src/_pytest/assertion/rewrite.py b/src/_pytest/assertion/rewrite.py index 2e606d1903a..62dac2bad0a 100644 --- a/src/_pytest/assertion/rewrite.py +++ b/src/_pytest/assertion/rewrite.py @@ -80,6 +80,7 @@ def __init__(self, config: Config) -> None: self._basenames_to_check_rewrite = {"conftest"} self._marked_for_rewrite_cache: dict[str, bool] = {} self._session_paths_checked = False + self.fn: str | None = None def set_session(self, session: Session | None) -> None: self.session = session @@ -126,7 +127,7 @@ def find_spec( ): return None else: - fn = spec.origin + self.fn = fn = spec.origin if not self._should_rewrite(name, fn, state): return None @@ -143,14 +144,11 @@ def create_module( ) -> types.ModuleType | None: return None # default behaviour is fine - def exec_module(self, module: types.ModuleType) -> None: - assert module.__spec__ is not None - assert module.__spec__.origin is not None - fn = Path(module.__spec__.origin) + def get_code(self, fullname: str) -> types.CodeType + assert self.fn is not None + fn = Path(self.fn) state = self.config.stash[assertstate_key] - self._rewritten_names[module.__name__] = fn - # The requested module looks like a test file, so rewrite it. This is # the most magical part of the process: load the source, rewrite the # asserts, and load the rewritten source. We also cache the rewritten @@ -183,7 +181,15 @@ def exec_module(self, module: types.ModuleType) -> None: self._writing_pyc = False else: state.trace(f"found cached rewritten pyc for {fn}") - exec(co, module.__dict__) + + return co + + def exec_module(self, module: types.ModuleType) -> None: + module_name = module.__name__ + + self._rewritten_names[module_name] = fn + + exec(self.get_code(module_name), module.__dict__) def _early_rewrite_bailout(self, name: str, state: AssertionState) -> bool: """A fast way to get out of rewriting modules. @@ -1213,4 +1219,4 @@ def get_cache_dir(file_path: Path) -> Path: return Path(sys.pycache_prefix) / Path(*file_path.parts[1:-1]) else: # classic pycache directory - return file_path.parent / "__pycache__" + return file_path.parent / "__pycache__" \ No newline at end of file From 873290003e7e7228ccbfd2f30cadab2a51c8b2e9 Mon Sep 17 00:00:00 2001 From: ShubhamDesai <42180509+ShubhamDesai@users.noreply.github.com> Date: Sat, 10 May 2025 14:04:14 -0400 Subject: [PATCH 2/2] Update src/_pytest/assertion/rewrite.py --- src/_pytest/assertion/rewrite.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/src/_pytest/assertion/rewrite.py b/src/_pytest/assertion/rewrite.py index 62dac2bad0a..a08b2b9834e 100644 --- a/src/_pytest/assertion/rewrite.py +++ b/src/_pytest/assertion/rewrite.py @@ -144,7 +144,7 @@ def create_module( ) -> types.ModuleType | None: return None # default behaviour is fine - def get_code(self, fullname: str) -> types.CodeType + def get_code(self, fullname: str) -> types.CodeType: assert self.fn is not None fn = Path(self.fn) state = self.config.stash[assertstate_key] @@ -183,12 +183,12 @@ def get_code(self, fullname: str) -> types.CodeType state.trace(f"found cached rewritten pyc for {fn}") return co - + def exec_module(self, module: types.ModuleType) -> None: module_name = module.__name__ - + self._rewritten_names[module_name] = fn - + exec(self.get_code(module_name), module.__dict__) def _early_rewrite_bailout(self, name: str, state: AssertionState) -> bool: