Skip to content
Merged
Show file tree
Hide file tree
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
22 changes: 20 additions & 2 deletions src/netshaper/core/mitm_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -209,24 +216,35 @@ 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 = {
"mitm_log_path": self._mitm_log_path,
}
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,
}
Expand Down
27 changes: 25 additions & 2 deletions src/netshaper/core/portal_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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)
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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
Expand All @@ -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(),
Expand Down
20 changes: 20 additions & 0 deletions src/netshaper/core/recovery_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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; "
Expand Down
139 changes: 134 additions & 5 deletions tests/test_managers.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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()

Expand All @@ -104,13 +109,48 @@ 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)

self.assertFalse(result)
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")

Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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()
Loading
Loading