diff --git a/hathor/sysctl/p2p/manager.py b/hathor/sysctl/p2p/manager.py index 2cfe291a6d..09e4ad407d 100644 --- a/hathor/sysctl/p2p/manager.py +++ b/hathor/sysctl/p2p/manager.py @@ -98,6 +98,11 @@ def __init__(self, connections: ConnectionsManager) -> None: self.get_enabled_sync_versions, self.set_enabled_sync_versions, ) + self.register( + 'kill_connection', + None, + self.set_kill_connection, + ) def set_force_sync_rotate(self) -> None: """Force a sync rotate.""" @@ -196,3 +201,16 @@ def _enable_sync_version(self, sync_version: SyncVersion) -> None: def _disable_sync_version(self, sync_version: SyncVersion) -> None: """Disable the given sync version.""" self.connections.disable_sync_version(sync_version) + + def set_kill_connection(self, peer_id: str, force: bool = False) -> None: + """Kill connection with peer_id or kill all connections if peer_id == '*'.""" + if peer_id == '*': + self.log.warn('Killing all connections') + self.connections.disconnect_all_peers(force=force) + return + + conn = self.connections.connected_peers.get(peer_id, None) + if conn is None: + self.log.warn('Killing connection', peer_id=peer_id) + raise SysctlException('peer-id is not connected') + conn.disconnect(force=force) diff --git a/hathor/sysctl/sysctl.py b/hathor/sysctl/sysctl.py index f9a805af84..1a045d0e36 100644 --- a/hathor/sysctl/sysctl.py +++ b/hathor/sysctl/sysctl.py @@ -15,12 +15,15 @@ from typing import Any, Callable, Iterator, NamedTuple, Optional from pydantic import validate_arguments +from structlog import get_logger from hathor.sysctl.exception import SysctlEntryNotFound, SysctlReadOnlyEntry, SysctlWriteOnlyEntry Getter = Callable[[], Any] Setter = Callable[..., None] +logger = get_logger() + class SysctlCommand(NamedTuple): getter: Optional[Getter] @@ -33,6 +36,7 @@ class Sysctl: def __init__(self) -> None: self._children: dict[str, 'Sysctl'] = {} self._commands: dict[str, SysctlCommand] = {} + self.log = logger.new() def put_child(self, path: str, sysctl: 'Sysctl') -> None: """Add a child to the tree.""" diff --git a/tests/sysctl/test_p2p.py b/tests/sysctl/test_p2p.py index 726e0d78ae..f8598b91de 100644 --- a/tests/sysctl/test_p2p.py +++ b/tests/sysctl/test_p2p.py @@ -151,6 +151,36 @@ def test_enabled_sync_versions(self): sysctl.set('enabled_sync_versions', ['v1']) self.assertEqual(sysctl.get('enabled_sync_versions'), ['v1']) + def test_kill_all_connections(self): + manager = self.create_peer() + p2p_manager = manager.connections + sysctl = ConnectionsManagerSysctl(p2p_manager) + + p2p_manager.disconnect_all_peers = MagicMock() + self.assertEqual(p2p_manager.disconnect_all_peers.call_count, 0) + sysctl.set('kill_connection', '*') + self.assertEqual(p2p_manager.disconnect_all_peers.call_count, 1) + + def test_kill_one_connection(self): + manager = self.create_peer() + p2p_manager = manager.connections + sysctl = ConnectionsManagerSysctl(p2p_manager) + + peer_id = 'my-peer-id' + conn = MagicMock() + p2p_manager.connected_peers[peer_id] = conn + self.assertEqual(conn.disconnect.call_count, 0) + sysctl.set('kill_connection', peer_id) + self.assertEqual(conn.disconnect.call_count, 1) + + def test_kill_connection_unknown_peer_id(self): + manager = self.create_peer() + p2p_manager = manager.connections + sysctl = ConnectionsManagerSysctl(p2p_manager) + + with self.assertRaises(SysctlException): + sysctl.set('kill_connection', 'unknown-peer-id') + class SyncV1RandomSimulatorTestCase(unittest.SyncV1Params, BaseRandomSimulatorTestCase): __test__ = True