From bd9f320eb2d670787a521c7ad0819205abad2733 Mon Sep 17 00:00:00 2001 From: Giorgio Salluzzo Date: Wed, 26 Aug 2026 07:49:18 +0200 Subject: [PATCH 1/7] Fix for nested Mocketizer's decorators. --- .pre-commit-config.yaml | 4 ++-- mocket/decorators/mocketizer.py | 23 +++++++++++++++-------- mocket/inject.py | 16 ++++++++++++++++ tests/test_mocket.py | 14 ++++++++++++++ tests/test_mode.py | 11 +++++++++++ 5 files changed, 58 insertions(+), 10 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index dfc6a580..5ee58ef7 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -15,12 +15,12 @@ repos: exclude: helm/ args: [ --unsafe ] - repo: https://github.com/charliermarsh/ruff-pre-commit - rev: "v0.15.18" + rev: "v0.16.4" hooks: - id: ruff args: [--fix, --exit-non-zero-on-fix] - id: ruff-format - repo: https://github.com/rstcheck/rstcheck - rev: v6.2.5 + rev: v6.3.0 hooks: - id: rstcheck diff --git a/mocket/decorators/mocketizer.py b/mocket/decorators/mocketizer.py index b067ffdf..68499307 100644 --- a/mocket/decorators/mocketizer.py +++ b/mocket/decorators/mocketizer.py @@ -32,13 +32,15 @@ def __init__( self.instance = instance self.truesocket_recording_dir = truesocket_recording_dir self.namespace = namespace or str(id(self)) - MocketMode.STRICT = strict_mode - if strict_mode: - MocketMode.STRICT_ALLOWED = strict_mode_allowed or [] - elif strict_mode_allowed: + if not strict_mode and strict_mode_allowed: raise ValueError( "Allowed locations are only accepted when STRICT mode is active." ) + self._previous_strict_mode = MocketMode.STRICT + self._previous_strict_mode_allowed = MocketMode.STRICT_ALLOWED + MocketMode.STRICT = strict_mode + if strict_mode: + MocketMode.STRICT_ALLOWED = strict_mode_allowed or [] def enter(self) -> None: """Enter the Mocketizer context (enable Mocket).""" @@ -60,10 +62,15 @@ def __enter__(self) -> Mocketizer: def exit(self) -> None: """Exit the Mocketizer context (disable Mocket).""" - if self.instance: - self.check_and_call("mocketize_teardown") - - Mocket.disable() + try: + if self.instance: + self.check_and_call("mocketize_teardown") + finally: + try: + Mocket.disable() + finally: + MocketMode.STRICT = self._previous_strict_mode + MocketMode.STRICT_ALLOWED = self._previous_strict_mode_allowed def __exit__(self, type: Any, value: Any, tb: Any) -> None: """Exit context manager. diff --git a/mocket/inject.py b/mocket/inject.py index e788a929..e129aaaf 100644 --- a/mocket/inject.py +++ b/mocket/inject.py @@ -11,6 +11,7 @@ import urllib3 _patches_restore: dict[tuple[ModuleType, str], Any] = {} +_enable_depth = 0 def _patch(module: ModuleType, name: str, patched_value: Any) -> None: @@ -39,6 +40,11 @@ def _restore(module: ModuleType, name: str) -> None: def enable() -> None: """Enable Mocket by patching socket, ssl, and urllib3 modules.""" + global _enable_depth + if _enable_depth > 0: + _enable_depth += 1 + return + from mocket.socket import ( MocketSocket, mock_create_connection, @@ -80,6 +86,8 @@ def enable() -> None: for (module, name), new_value in patches.items(): _patch(module, name, new_value) + _enable_depth += 1 + with contextlib.suppress(ImportError): from urllib3.contrib.pyopenssl import extract_from_urllib3 @@ -88,6 +96,14 @@ def enable() -> None: def disable() -> None: """Disable Mocket by restoring all patched modules.""" + global _enable_depth + if _enable_depth == 0: + return + + _enable_depth -= 1 + if _enable_depth > 0: + return + for module, name in list(_patches_restore.keys()): _restore(module, name) diff --git a/tests/test_mocket.py b/tests/test_mocket.py index 8810a5b9..0097faaf 100644 --- a/tests/test_mocket.py +++ b/tests/test_mocket.py @@ -10,6 +10,7 @@ from mocket import Mocket, MocketEntry, Mocketizer, mocketize from mocket.compat import encode_to_bytes +from mocket.mode import MocketMode class MocketTestCase(TestCase): @@ -221,6 +222,19 @@ def test_patch( assert os.getcwd() == "foo" +def test_mocketize_twice_nested(): + original_socket = socket.socket + original_strict_mode = MocketMode.STRICT + + with Mocketizer(strict_mode=True), Mocketizer(strict_mode=True): + pass + + assert socket.socket is original_socket + assert MocketMode.STRICT is original_strict_mode + url = "http://httpbin.local/ip" + assert httpx.get(url).status_code == 200 + + @pytest.mark.skipif(not psutil.POSIX, reason="Uses a POSIX-only API to test") @pytest.mark.skipif('os.getenv("SKIP_TRUE_HTTP", False)') @pytest.mark.asyncio diff --git a/tests/test_mode.py b/tests/test_mode.py index bfdb2a79..241cf834 100644 --- a/tests/test_mode.py +++ b/tests/test_mode.py @@ -71,3 +71,14 @@ def test_strict_mode_allowed_or_not(strict_mode_on): with Mocketizer(strict_mode=strict_mode_on): assert MocketMode.is_allowed("foobar.com") is not strict_mode_on assert MocketMode.is_allowed(("foobar.com", 443)) is not strict_mode_on + + +def test_mocketize_strict_mode_does_not_leak_after_outer_context(): + with Mocketizer(strict_mode=False): + + @mocketize(strict_mode=True) + def strict_test(): + assert MocketMode.STRICT is True + + strict_test() + assert MocketMode.STRICT is False From 8ad1c61f5cdf8fa49d168a3ed533cff5afecaf22 Mon Sep 17 00:00:00 2001 From: Giorgio Salluzzo Date: Wed, 26 Aug 2026 07:58:05 +0200 Subject: [PATCH 2/7] Coverage back to 100%. --- mocket/mocket.py | 2 +- tests/test_socket.py | 30 +++++++++++++++++++++++++++++- 2 files changed, 30 insertions(+), 2 deletions(-) diff --git a/mocket/mocket.py b/mocket/mocket.py index 9baa2ac2..353c553f 100644 --- a/mocket/mocket.py +++ b/mocket/mocket.py @@ -14,7 +14,7 @@ # NOTE this is here for backwards-compat to keep old import-paths working # from mocket.socket import MocketSocket as MocketSocket -if TYPE_CHECKING: +if TYPE_CHECKING: # pragma: no cover from mocket.entry import MocketEntry from mocket.types import Address diff --git a/tests/test_socket.py b/tests/test_socket.py index 31e0d63a..5014077b 100644 --- a/tests/test_socket.py +++ b/tests/test_socket.py @@ -9,8 +9,9 @@ from mocket import Mocket, MocketEntry, Mocketizer, mocketize from mocket.mockhttp import Entry from mocket.socket import MocketSocket -from mocket.ssl.context import MocketSSLContext +from mocket.ssl.context import MocketSSLContext, mock_wrap_socket from mocket.ssl.socket import MocketSSLSocket +from mocket.urllib3 import mock_match_hostname @pytest.mark.parametrize("blocking", (False, True)) @@ -163,6 +164,33 @@ def test_wrap_bio_preserves_empty_server_hostname_on_getpeercert(monkeypatch): assert ssl_obj._address == ("", 443) +def test_wrap_bio_with_invalid_mocket_address(monkeypatch): + monkeypatch.setattr(Mocket, "_address", "invalid-address") + ssl_obj = MocketSSLContext().wrap_bio( + incoming=None, + outgoing=None, + server_hostname=None, + ) + + assert ssl_obj._host is None + assert ssl_obj._port is None + + +def test_mock_wrap_socket_delegates_to_context(monkeypatch): + expected = MocketSSLSocket() + + def fake_wrap_socket(self, sock, *args, **kwargs): + return expected + + monkeypatch.setattr(MocketSSLContext, "wrap_socket", fake_wrap_socket) + + assert mock_wrap_socket(MocketSocket()) is expected + + +def test_mock_match_hostname_returns_none(): + assert mock_match_hostname("example.org", object()) is None + + def test_getpeercert_does_not_overwrite_empty_host_when_port_missing(monkeypatch): monkeypatch.setattr(Mocket, "_address", ("httpbin.local", 443)) ssl_obj = MocketSSLSocket() From dc67b9dd1cfeabaa50f87184b5990b0b575700ca Mon Sep 17 00:00:00 2001 From: Giorgio Salluzzo Date: Wed, 26 Aug 2026 07:59:07 +0200 Subject: [PATCH 3/7] Potential fix for pull request finding Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- tests/test_mocket.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/test_mocket.py b/tests/test_mocket.py index 0097faaf..3e4de357 100644 --- a/tests/test_mocket.py +++ b/tests/test_mocket.py @@ -222,6 +222,7 @@ def test_patch( assert os.getcwd() == "foo" +@pytest.mark.skipif('os.getenv("SKIP_TRUE_HTTP", False)') def test_mocketize_twice_nested(): original_socket = socket.socket original_strict_mode = MocketMode.STRICT @@ -232,7 +233,7 @@ def test_mocketize_twice_nested(): assert socket.socket is original_socket assert MocketMode.STRICT is original_strict_mode url = "http://httpbin.local/ip" - assert httpx.get(url).status_code == 200 + assert httpx.get(url, timeout=5.0).status_code == 200 @pytest.mark.skipif(not psutil.POSIX, reason="Uses a POSIX-only API to test") From eee95d55955cb6a717bdd1cad2e34707f40075d8 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 26 Aug 2026 06:00:38 +0000 Subject: [PATCH 4/7] Protect _enable_depth with a threading.Lock in enable/disable Co-authored-by: mindflayer <527325+mindflayer@users.noreply.github.com> --- mocket/inject.py | 27 ++++++++++++++++----------- 1 file changed, 16 insertions(+), 11 deletions(-) diff --git a/mocket/inject.py b/mocket/inject.py index e129aaaf..478c654d 100644 --- a/mocket/inject.py +++ b/mocket/inject.py @@ -5,6 +5,7 @@ import contextlib import socket import ssl +import threading from types import ModuleType from typing import Any @@ -12,6 +13,7 @@ _patches_restore: dict[tuple[ModuleType, str], Any] = {} _enable_depth = 0 +_enable_lock = threading.Lock() def _patch(module: ModuleType, name: str, patched_value: Any) -> None: @@ -41,9 +43,6 @@ def _restore(module: ModuleType, name: str) -> None: def enable() -> None: """Enable Mocket by patching socket, ssl, and urllib3 modules.""" global _enable_depth - if _enable_depth > 0: - _enable_depth += 1 - return from mocket.socket import ( MocketSocket, @@ -83,10 +82,15 @@ def enable() -> None: (urllib3.util.ssl_, "wrap_socket"): mock_urllib3_ssl_wrap_socket, # urllib3 < 2 } - for (module, name), new_value in patches.items(): - _patch(module, name, new_value) + with _enable_lock: + if _enable_depth > 0: + _enable_depth += 1 + return + + for (module, name), new_value in patches.items(): + _patch(module, name, new_value) - _enable_depth += 1 + _enable_depth += 1 with contextlib.suppress(ImportError): from urllib3.contrib.pyopenssl import extract_from_urllib3 @@ -97,12 +101,13 @@ def enable() -> None: def disable() -> None: """Disable Mocket by restoring all patched modules.""" global _enable_depth - if _enable_depth == 0: - return + with _enable_lock: + if _enable_depth == 0: + return - _enable_depth -= 1 - if _enable_depth > 0: - return + _enable_depth -= 1 + if _enable_depth > 0: + return for module, name in list(_patches_restore.keys()): _restore(module, name) From fb0e28e5da9f3d168dc02b5a8f0ed589760c9316 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 26 Aug 2026 06:03:29 +0000 Subject: [PATCH 5/7] Move _restore calls inside _enable_lock in disable() Co-authored-by: mindflayer <527325+mindflayer@users.noreply.github.com> --- mocket/inject.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/mocket/inject.py b/mocket/inject.py index 478c654d..c5c38cc2 100644 --- a/mocket/inject.py +++ b/mocket/inject.py @@ -109,8 +109,8 @@ def disable() -> None: if _enable_depth > 0: return - for module, name in list(_patches_restore.keys()): - _restore(module, name) + for module, name in list(_patches_restore.keys()): + _restore(module, name) with contextlib.suppress(ImportError): from urllib3.contrib.pyopenssl import inject_into_urllib3 From ee6666d29e91002856bd0bbea38e0255692b3d24 Mon Sep 17 00:00:00 2001 From: Giorgio Salluzzo Date: Wed, 26 Aug 2026 08:04:00 +0200 Subject: [PATCH 6/7] Coverage back to 100%. --- tests/test_inject.py | 43 +++++++++++++++++++++++++++++++++++++++++ tests/test_recording.py | 43 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 86 insertions(+) create mode 100644 tests/test_inject.py create mode 100644 tests/test_recording.py diff --git a/tests/test_inject.py b/tests/test_inject.py new file mode 100644 index 00000000..2d57257e --- /dev/null +++ b/tests/test_inject.py @@ -0,0 +1,43 @@ +import sys +import types + +from mocket import inject + + +def test_disable_calls_pyopenssl_inject_when_available(monkeypatch): + calls: list[str] = [] + pyopenssl_module = types.ModuleType("urllib3.contrib.pyopenssl") + + def fake_inject_into_urllib3(): + calls.append("called") + + pyopenssl_module.inject_into_urllib3 = fake_inject_into_urllib3 + monkeypatch.setitem(sys.modules, "urllib3.contrib.pyopenssl", pyopenssl_module) + monkeypatch.setattr(inject, "_patches_restore", {}) + monkeypatch.setattr(inject, "_enable_depth", 1) + + inject.disable() + + assert calls == ["called"] + + +def test_enable_calls_pyopenssl_extract_when_available(monkeypatch): + calls: list[str] = [] + pyopenssl_module = types.ModuleType("urllib3.contrib.pyopenssl") + + def fake_extract_from_urllib3(): + calls.append("called") + + def fake_inject_into_urllib3(): + pass + + pyopenssl_module.extract_from_urllib3 = fake_extract_from_urllib3 + pyopenssl_module.inject_into_urllib3 = fake_inject_into_urllib3 + monkeypatch.setitem(sys.modules, "urllib3.contrib.pyopenssl", pyopenssl_module) + monkeypatch.setattr(inject, "_patches_restore", {}) + monkeypatch.setattr(inject, "_enable_depth", 0) + + inject.enable() + inject.disable() + + assert calls == ["called"] diff --git a/tests/test_recording.py b/tests/test_recording.py new file mode 100644 index 00000000..92c5e47d --- /dev/null +++ b/tests/test_recording.py @@ -0,0 +1,43 @@ +from mocket.recording import MocketRecord, MocketRecordStorage, _hash_request_fallback + + +def test_get_records_returns_all_records_for_address(tmp_path): + storage = MocketRecordStorage(directory=tmp_path, namespace="recording-get-records") + address = ("example.org", 80) + signature = "signature" + storage._records[address][signature] = MocketRecord( + host=address[0], + port=address[1], + request=b"GET / HTTP/1.1\r\nHost: example.org\r\n\r\n", + response=b"HTTP/1.1 200 OK\r\n\r\nok", + ) + + records = storage.get_records(address) + + assert len(records) == 1 + assert records[0].response == b"HTTP/1.1 200 OK\r\n\r\nok" + + +def test_put_record_updates_fallback_signature_without_saving(tmp_path): + storage = MocketRecordStorage( + directory=tmp_path, namespace="recording-put-record-fallback" + ) + address = ("example.org", 80) + request = b"GET / HTTP/1.1\r\nHost: example.org\r\n\r\n" + fallback_signature = _hash_request_fallback(request) + + storage._records[address][fallback_signature] = MocketRecord( + host=address[0], + port=address[1], + request=request, + response=b"HTTP/1.1 200 OK\r\n\r\nold", + ) + + storage.put_record( + address=address, + request=request, + response=b"HTTP/1.1 200 OK\r\n\r\nnew", + ) + + assert storage._records[address][fallback_signature].response.endswith(b"new") + assert not storage.file.exists() From 60a36fc2f270a9d99824d6943b5c2ec3142a68da Mon Sep 17 00:00:00 2001 From: Giorgio Salluzzo Date: Wed, 26 Aug 2026 08:08:48 +0200 Subject: [PATCH 7/7] Moving new test to where it belongs. --- tests/test_mocket.py | 15 --------------- tests/test_mode.py | 17 +++++++++++++++++ 2 files changed, 17 insertions(+), 15 deletions(-) diff --git a/tests/test_mocket.py b/tests/test_mocket.py index 3e4de357..8810a5b9 100644 --- a/tests/test_mocket.py +++ b/tests/test_mocket.py @@ -10,7 +10,6 @@ from mocket import Mocket, MocketEntry, Mocketizer, mocketize from mocket.compat import encode_to_bytes -from mocket.mode import MocketMode class MocketTestCase(TestCase): @@ -222,20 +221,6 @@ def test_patch( assert os.getcwd() == "foo" -@pytest.mark.skipif('os.getenv("SKIP_TRUE_HTTP", False)') -def test_mocketize_twice_nested(): - original_socket = socket.socket - original_strict_mode = MocketMode.STRICT - - with Mocketizer(strict_mode=True), Mocketizer(strict_mode=True): - pass - - assert socket.socket is original_socket - assert MocketMode.STRICT is original_strict_mode - url = "http://httpbin.local/ip" - assert httpx.get(url, timeout=5.0).status_code == 200 - - @pytest.mark.skipif(not psutil.POSIX, reason="Uses a POSIX-only API to test") @pytest.mark.skipif('os.getenv("SKIP_TRUE_HTTP", False)') @pytest.mark.asyncio diff --git a/tests/test_mode.py b/tests/test_mode.py index 241cf834..b885f7b0 100644 --- a/tests/test_mode.py +++ b/tests/test_mode.py @@ -1,3 +1,6 @@ +import socket + +import httpx import pytest import requests @@ -82,3 +85,17 @@ def strict_test(): strict_test() assert MocketMode.STRICT is False + + +@pytest.mark.skipif('os.getenv("SKIP_TRUE_HTTP", False)') +def test_mocketize_twice_nested(): + original_socket = socket.socket + original_strict_mode = MocketMode.STRICT + + with Mocketizer(strict_mode=True), Mocketizer(strict_mode=True): + pass + + assert socket.socket is original_socket + assert MocketMode.STRICT is original_strict_mode + url = "http://httpbin.local/ip" + assert httpx.get(url, timeout=5.0).status_code == 200