From b5ebced4e15f01874941bb14ee79b6895e5abd92 Mon Sep 17 00:00:00 2001 From: Kyrol Date: Wed, 8 Jul 2026 16:01:05 +0800 Subject: [PATCH] Harden managed service identity recovery --- src/netshaper/core/mitm_manager.py | 22 +++- src/netshaper/core/portal_manager.py | 27 ++++- src/netshaper/core/recovery_manager.py | 20 ++++ tests/test_managers.py | 139 ++++++++++++++++++++++++- tests/test_orchestrator.py | 19 +++- tests/test_portal_manager.py | 75 +++++++++++-- 6 files changed, 285 insertions(+), 17 deletions(-) diff --git a/src/netshaper/core/mitm_manager.py b/src/netshaper/core/mitm_manager.py index b9311ae..e45f4e7 100644 --- a/src/netshaper/core/mitm_manager.py +++ b/src/netshaper/core/mitm_manager.py @@ -45,6 +45,7 @@ def __init__( self.own_ip = own_ip self._mitm_proc: Optional[subprocess.Popen[Any]] = None self._mitm_command: Optional[list[str]] = None + self._mitm_process_identity: Optional[dict[str, object]] = None self._mitm_log_path: Optional[str] = None self._mitm_log_handle: Optional[object] = None self._journal = journal @@ -76,6 +77,7 @@ def _close_log(self) -> None: def _clear_completed_process(self) -> None: self._mitm_proc = None self._mitm_command = None + self._mitm_process_identity = None self._close_log() @staticmethod @@ -125,12 +127,17 @@ def launch(self, port: int = 8088, web_port: int = 8083) -> bool: "--set", f"web_port={web_port}", ] - self._mitm_proc = subprocess.Popen( # nosec B603 B607 + process = subprocess.Popen( # nosec B603 B607 command, stdout=log_handle or subprocess.DEVNULL, stderr=subprocess.STDOUT if log_handle else subprocess.DEVNULL, ) + self._mitm_proc = process self._mitm_command = list(command) + if not self._capture_mitm_process_identity(process): + log.error("Could not establish recoverable mitmproxy process identity") + self.terminate() + return False if not self._journal_state(): log.error("Refusing to run mitmproxy without recovery state") self.terminate() @@ -209,12 +216,20 @@ def terminate(self) -> bool: if ok: self._mitm_proc = None self._mitm_command = None + self._mitm_process_identity = None self._close_log() if not self._journal_state(): ok = False return ok + def _capture_mitm_process_identity(self, process: subprocess.Popen[Any]) -> bool: + identity = process_owner_metadata(process.pid) + if identity.get("process_create_time") is None: + return False + self._mitm_process_identity = dict(identity) + return True + def get_state_for_persistence(self) -> dict: """Get mitmproxy state for persistence.""" state = { @@ -222,11 +237,14 @@ def get_state_for_persistence(self) -> dict: } proc = self._mitm_proc if proc is not None and proc.poll() is None: + if self._mitm_process_identity is None: + log.error("Live mitmproxy process has no recoverable identity") + return state command = list(self._mitm_command or getattr(proc, "args", []) or []) state.update( { "service": "mitmproxy", - **process_owner_metadata(proc.pid), + **self._mitm_process_identity, "executable": command[0] if command else None, "argv": command, } diff --git a/src/netshaper/core/portal_manager.py b/src/netshaper/core/portal_manager.py index 4f8cc89..17ac276 100644 --- a/src/netshaper/core/portal_manager.py +++ b/src/netshaper/core/portal_manager.py @@ -45,6 +45,7 @@ def __init__( self.authorized_cidrs = tuple(str(network) for network in authorized_cidrs) self.process: Optional[subprocess.Popen[Any]] = None self._command: Optional[list[str]] = None + self._process_identity: Optional[dict[str, object]] = None self._health_token: Optional[str] = None self._journal = journal @@ -68,6 +69,12 @@ def start(self, portal_config: PortalConfig) -> bool: return True if self.process and self.process.poll() is None: + if self._process_identity is None and not self._capture_process_identity( + self.process, + ): + log.error("Could not establish recoverable portal process identity") + self.stop() + return False log.debug("Waiting for existing netshaper-portal child") else: dns_claimed = check_local_port(self.host_ip, 53, socket.SOCK_DGRAM) @@ -105,15 +112,20 @@ def start(self, portal_config: PortalConfig) -> bool: cmd.append("--hsts-idn-demo") try: - self.process = subprocess.Popen( # nosec B603 + process = subprocess.Popen( # nosec B603 cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, ) + self.process = process self._command = list(cmd) except OSError as exc: log.error(f"portal launch failed: {exc}") return False + if not self._capture_process_identity(process): + log.error("Could not establish recoverable portal process identity") + self.stop() + return False if not self._journal_state(): log.error("Refusing to run netshaper-portal without recovery state") self.stop() @@ -190,10 +202,18 @@ def stop(self) -> bool: if ok: self.process = None self._command = None + self._process_identity = None if not self._journal_state(): ok = False return ok + def _capture_process_identity(self, process: subprocess.Popen[Any]) -> bool: + identity = process_owner_metadata(process.pid) + if identity.get("process_create_time") is None: + return False + self._process_identity = dict(identity) + return True + def _journal_state(self) -> bool: if self._journal is None: return True @@ -207,11 +227,14 @@ def get_state_for_persistence(self) -> dict[str, object]: process = self.process if process is None or process.poll() is not None: return {} + if self._process_identity is None: + log.error("Live portal process has no recoverable identity") + return {} command = list(self._command or getattr(process, "args", []) or []) executable = command[0] if command else None return { "service": "portal", - **process_owner_metadata(process.pid), + **self._process_identity, "executable": executable, "argv": command, "ownership_token": self.health_token(), diff --git a/src/netshaper/core/recovery_manager.py b/src/netshaper/core/recovery_manager.py index c58335b..6b24758 100644 --- a/src/netshaper/core/recovery_manager.py +++ b/src/netshaper/core/recovery_manager.py @@ -354,6 +354,26 @@ def _cleanup_managed_service( ) return False + try: + expected_create_time = float(record["process_create_time"]) + actual_create_time = process.create_time() + except psutil.NoSuchProcess: + return True + except Exception as exc: + log.error( + "[Recovery] Could not verify managed service %s process birth: %s", + service_name, + exc, + ) + return False + if actual_create_time != expected_create_time: + log.error( + "[Recovery] Managed service %s process birth mismatch; " + "leaving it in place", + service_name, + ) + return False + if not cls._process_matches_service_record(service_name, process, record): log.error( "[Recovery] Managed service %s process identity mismatch; " diff --git a/tests/test_managers.py b/tests/test_managers.py index fbc6486..c594422 100644 --- a/tests/test_managers.py +++ b/tests/test_managers.py @@ -1,10 +1,14 @@ import json import os +import subprocess +import sys import tempfile import unittest from ipaddress import IPv4Network from unittest import mock +import psutil + from netshaper import config from netshaper.core.authorization import AuthorizationError, AuthorizationPolicy from netshaper.core.firewall_manager import FirewallManager @@ -75,14 +79,15 @@ def test_mitm_manager_state_includes_owned_process(self): "8088", ] manager._mitm_log_path = "/run/netshaper/NS-TEST/mitmproxy.log" + manager._mitm_process_identity = { + "pid": 4321, + "process_create_time": 123.0, + "created_at": 456.0, + } with mock.patch( "netshaper.core.mitm_manager.process_owner_metadata", - return_value={ - "pid": 4321, - "process_create_time": 123.0, - "created_at": 456.0, - }, + side_effect=AssertionError("identity should be captured once"), ): state = manager.get_state_for_persistence() @@ -104,6 +109,14 @@ def test_mitm_manager_stops_process_when_launch_journal_fails(self): return_value=False), \ mock.patch("netshaper.core.mitm_manager.subprocess.Popen", return_value=process), \ + mock.patch( + "netshaper.core.mitm_manager.process_owner_metadata", + return_value={ + "pid": 4321, + "process_create_time": 123.0, + "created_at": 456.0, + }, + ), \ mock.patch("netshaper.core.mitm_manager.log"): result = manager.launch(port=8088, web_port=8083) @@ -111,6 +124,33 @@ def test_mitm_manager_stops_process_when_launch_journal_fails(self): process.terminate.assert_called_once() self.assertIsNone(manager._mitm_proc) + def test_mitm_manager_stops_process_when_identity_capture_fails(self): + process = mock.Mock() + process.pid = 4321 + process.poll.side_effect = [None, 0] + manager = MitmProxyManager("127.0.0.1") + + with tempfile.TemporaryDirectory() as tmp, \ + mock.patch.object(config, "STATE_DIR", tmp), \ + mock.patch.object(config, "DRY_RUN", False), \ + mock.patch("netshaper.core.mitm_manager.check_local_port", + return_value=False), \ + mock.patch("netshaper.core.mitm_manager.subprocess.Popen", + return_value=process), \ + mock.patch( + "netshaper.core.mitm_manager.process_owner_metadata", + return_value={ + "pid": 4321, + "process_create_time": None, + "created_at": 456.0, + }, + ): + result = manager.launch(port=8088, web_port=8083) + + self.assertFalse(result) + process.terminate.assert_called_once() + self.assertIsNone(manager._mitm_proc) + def test_mitm_manager_refuses_existing_listener(self): manager = MitmProxyManager("127.0.0.1") @@ -233,6 +273,7 @@ def test_recovery_keeps_unknown_active_plugin_manifest(self): def test_recovery_terminates_verified_managed_service(self): process = mock.Mock() + process.create_time.return_value = 1.0 process.cmdline.return_value = [ "/usr/bin/python3", "-m", @@ -265,6 +306,7 @@ def test_recovery_terminates_verified_managed_service(self): def test_recovery_kills_verified_managed_service_after_timeout(self): process = mock.Mock() + process.create_time.return_value = 1.0 process.cmdline.return_value = [ "mitmweb", "--mode", @@ -298,6 +340,7 @@ def test_recovery_kills_verified_managed_service_after_timeout(self): def test_recovery_refuses_managed_service_identity_mismatch(self): process = mock.Mock() + process.create_time.return_value = 1.0 process.cmdline.return_value = ["python", "-m", "other.portal"] record = { "service": "portal", @@ -327,6 +370,39 @@ def test_recovery_refuses_managed_service_identity_mismatch(self): process.terminate.assert_not_called() + def test_recovery_refuses_managed_service_birth_mismatch(self): + process = mock.Mock() + process.create_time.return_value = 2.0 + process.cmdline.return_value = [ + "/usr/bin/python3", + "-m", + "netshaper.portal", + "--health-token", + "token", + ] + record = { + "service": "portal", + "pid": 1234, + "process_create_time": 1.0, + "executable": "/usr/bin/python3", + "argv": process.cmdline.return_value, + "ownership_token": "token", + } + + with mock.patch( + "netshaper.core.recovery_manager.owner_status", + return_value=OwnerStatus.LIVE, + ), mock.patch( + "netshaper.core.recovery_manager.psutil.Process", + return_value=process, + ): + self.assertFalse( + RecoveryManager._cleanup_managed_service("portal", record) + ) + + process.cmdline.assert_not_called() + process.terminate.assert_not_called() + def test_recovery_fails_closed_for_unknown_managed_service_owner(self): record = { "service": "mitmproxy", @@ -365,6 +441,59 @@ def test_recovery_treats_stale_managed_service_as_clean(self): RecoveryManager._cleanup_managed_service("mitmproxy", record) ) + def test_recovery_terminates_live_harmless_managed_child_only(self): + token = "test-token" + managed_cmd = [ + sys.executable, + "-c", + "import time; time.sleep(60)", + "netshaper.portal", + "--health-token", + token, + ] + unrelated_cmd = [ + sys.executable, + "-c", + "import time; time.sleep(60)", + "unrelated-process", + ] + managed = subprocess.Popen( # nosec B603 + managed_cmd, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + unrelated = subprocess.Popen( # nosec B603 + unrelated_cmd, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + try: + managed_process = psutil.Process(managed.pid) + record = { + "service": "portal", + "pid": managed.pid, + "process_create_time": managed_process.create_time(), + "executable": sys.executable, + "argv": managed_cmd, + "ownership_token": token, + } + + self.assertTrue( + RecoveryManager._cleanup_managed_service("portal", record) + ) + + managed.wait(timeout=5) + self.assertIsNone(unrelated.poll()) + finally: + for child in (managed, unrelated): + if child.poll() is None: + child.terminate() + try: + child.wait(timeout=5) + except subprocess.TimeoutExpired: + child.kill() + child.wait(timeout=5) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_orchestrator.py b/tests/test_orchestrator.py index ed8f626..4930b99 100644 --- a/tests/test_orchestrator.py +++ b/tests/test_orchestrator.py @@ -374,6 +374,7 @@ def test_fake_server_launch_wires_spoof_mode_dnssec_and_allowlist(self): ns._fake_server_proc = None ns._fake_server_health_token = "test-health-token" process = mock.Mock() + process.pid = 1234 process.poll.return_value = None with mock.patch( @@ -385,7 +386,14 @@ def test_fake_server_launch_wires_spoof_mode_dnssec_and_allowlist(self): ), mock.patch( "netshaper.core.portal_manager.subprocess.Popen", return_value=process, - ) as popen: + ) as popen, mock.patch( + "netshaper.core.portal_manager.process_owner_metadata", + return_value={ + "pid": 1234, + "process_create_time": 456.0, + "created_at": 789.0, + }, + ): self.assertTrue( ns.launch_fake_server( dnssec_mode="nxdomain", @@ -515,6 +523,7 @@ def test_launch_mitmproxy_reaps_process_when_readiness_fails(self): ns.session_id = "NS-TEST" ns._mitm_log_path = None proc = mock.Mock() + proc.pid = 4321 proc.poll.side_effect = [None, None, 0] with tempfile.TemporaryDirectory() as tmp, \ @@ -526,6 +535,14 @@ def test_launch_mitmproxy_reaps_process_when_readiness_fails(self): return_value=False), \ mock.patch("netshaper.core.mitm_manager.subprocess.Popen", return_value=proc), \ + mock.patch( + "netshaper.core.mitm_manager.process_owner_metadata", + return_value={ + "pid": 4321, + "process_create_time": 123.0, + "created_at": 456.0, + }, + ), \ mock.patch("netshaper.core.mitm_manager.time.sleep"), \ mock.patch("netshaper.core.mitm_manager.log"): result = ns.launch_mitmproxy() diff --git a/tests/test_portal_manager.py b/tests/test_portal_manager.py index c7a9c7b..ed07951 100644 --- a/tests/test_portal_manager.py +++ b/tests/test_portal_manager.py @@ -48,15 +48,16 @@ def test_state_includes_owned_portal_process(self): "--health-token", "token", ] + manager._process_identity = { + "pid": 1234, + "process_create_time": 456.0, + "created_at": 789.0, + } manager._health_token = "token" with mock.patch( "netshaper.core.portal_manager.process_owner_metadata", - return_value={ - "pid": 1234, - "process_create_time": 456.0, - "created_at": 789.0, - }, + side_effect=AssertionError("identity should be captured once"), ): state = manager.get_state_for_persistence() @@ -68,13 +69,22 @@ def test_state_includes_owned_portal_process(self): def test_start_waits_for_existing_child(self): manager = PortalManager("192.0.2.10", ["192.0.2.0/24"]) manager.process = mock.Mock() + manager.process.pid = 1234 manager.process.poll.return_value = None with mock.patch.object(manager, "ready", side_effect=[False, True]), \ mock.patch("netshaper.core.portal_manager.check_local_port" ) as port_check, \ mock.patch("netshaper.core.portal_manager.subprocess.Popen" - ) as popen: + ) as popen, \ + mock.patch( + "netshaper.core.portal_manager.process_owner_metadata", + return_value={ + "pid": 1234, + "process_create_time": 456.0, + "created_at": 789.0, + }, + ): self.assertTrue(manager.start(PortalConfig())) port_check.assert_not_called() @@ -83,6 +93,7 @@ def test_start_waits_for_existing_child(self): def test_start_launches_public_portal_module(self): manager = PortalManager("192.0.2.10", ["192.0.2.0/24"]) process = mock.Mock() + process.pid = 1234 process.poll.return_value = None with mock.patch.object(manager, "ready", side_effect=[False, True]), \ @@ -91,7 +102,15 @@ def test_start_launches_public_portal_module(self): mock.patch("netshaper.core.portal_manager.check_local_port", side_effect=[False, False]), \ mock.patch("netshaper.core.portal_manager.subprocess.Popen", - return_value=process) as popen: + return_value=process) as popen, \ + mock.patch( + "netshaper.core.portal_manager.process_owner_metadata", + return_value={ + "pid": 1234, + "process_create_time": 456.0, + "created_at": 789.0, + }, + ): self.assertTrue( manager.start( PortalConfig( @@ -111,6 +130,30 @@ def test_start_launches_public_portal_module(self): self.assertIn("192.0.2.0/24", command) self.assertIn("192.0.2.10/32", command) + def test_start_stops_child_when_identity_capture_fails(self): + manager = PortalManager("192.0.2.10", ["192.0.2.0/24"]) + process = mock.Mock() + process.pid = 1234 + process.poll.side_effect = [None, 0] + + with mock.patch.object(manager, "ready", return_value=False), \ + mock.patch("netshaper.core.portal_manager.check_local_port", + side_effect=[False, False]), \ + mock.patch("netshaper.core.portal_manager.subprocess.Popen", + return_value=process), \ + mock.patch( + "netshaper.core.portal_manager.process_owner_metadata", + return_value={ + "pid": 1234, + "process_create_time": None, + "created_at": 789.0, + }, + ): + self.assertFalse(manager.start(PortalConfig())) + + process.terminate.assert_called_once() + self.assertIsNone(manager.process) + def test_start_returns_false_when_popen_fails(self): manager = PortalManager("192.0.2.10", ["192.0.2.0/24"]) @@ -124,6 +167,7 @@ def test_start_returns_false_when_popen_fails(self): def test_start_returns_false_when_child_exits_during_startup(self): manager = PortalManager("192.0.2.10", ["192.0.2.0/24"]) process = mock.Mock() + process.pid = 1234 process.poll.return_value = 2 process.returncode = 2 @@ -132,6 +176,14 @@ def test_start_returns_false_when_child_exits_during_startup(self): side_effect=[False, False]), \ mock.patch("netshaper.core.portal_manager.subprocess.Popen", return_value=process), \ + mock.patch( + "netshaper.core.portal_manager.process_owner_metadata", + return_value={ + "pid": 1234, + "process_create_time": 456.0, + "created_at": 789.0, + }, + ), \ mock.patch.object(manager, "stop", wraps=manager.stop) as stop_mock: self.assertFalse(manager.start(PortalConfig())) @@ -140,6 +192,7 @@ def test_start_returns_false_when_child_exits_during_startup(self): def test_start_times_out_and_stops_child(self): manager = PortalManager("192.0.2.10", ["192.0.2.0/24"]) process = mock.Mock() + process.pid = 1234 process.poll.return_value = None with mock.patch.object(manager, "ready", return_value=False), \ @@ -147,6 +200,14 @@ def test_start_times_out_and_stops_child(self): side_effect=[False, False]), \ mock.patch("netshaper.core.portal_manager.subprocess.Popen", return_value=process), \ + mock.patch( + "netshaper.core.portal_manager.process_owner_metadata", + return_value={ + "pid": 1234, + "process_create_time": 456.0, + "created_at": 789.0, + }, + ), \ mock.patch("netshaper.core.portal_manager.time.sleep"), \ mock.patch.object(manager, "stop", return_value=True) as stop_mock: self.assertFalse(manager.start(PortalConfig()))