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
18 changes: 18 additions & 0 deletions hathor/sysctl/p2p/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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)
4 changes: 4 additions & 0 deletions hathor/sysctl/sysctl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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."""
Expand Down
30 changes: 30 additions & 0 deletions tests/sysctl/test_p2p.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down