From 03933c2b3a838e6ee8f80b81886e91092355584e Mon Sep 17 00:00:00 2001 From: Arnaud Date: Fri, 22 May 2026 22:54:07 +0400 Subject: [PATCH] Add hole punching --- storage/nat.nim | 80 ++++++++++++++++++++++++++++++++++++++- storage/storage.nim | 2 + tests/storage/testnat.nim | 35 ++++++++++++++++- 3 files changed, 113 insertions(+), 4 deletions(-) diff --git a/storage/nat.nim b/storage/nat.nim index fcf07234..b1729713 100644 --- a/storage/nat.nim +++ b/storage/nat.nim @@ -16,6 +16,10 @@ import pkg/chronicles import pkg/libp2p import pkg/libp2p/services/autorelayservice import pkg/libp2p/protocols/connectivity/autonatv2/service +import pkg/libp2p/protocols/connectivity/relay/relay as relayProtocol +import pkg/libp2p/protocols/connectivity/dcutr/client as dcutrClientModule +import pkg/libp2p/protocols/connectivity/dcutr/server as dcutrServerModule +import pkg/libp2p/wire import ./utils import ./utils/natutils @@ -60,8 +64,7 @@ method mapNatPorts*( if not m.plumInitialized: # 5s matches the old NatPortMappingTimeout used with miniupnpc/libnatpmp. let plumLogLevel = - if getEnv("DEBUG") == "1": PLUM_LOG_LEVEL_VERBOSE - else: PLUM_LOG_LEVEL_NONE + if getEnv("DEBUG") == "1": PLUM_LOG_LEVEL_VERBOSE else: PLUM_LOG_LEVEL_NONE let res = init( logLevel = plumLogLevel, discoverTimeout = m.discoverTimeout, @@ -229,3 +232,76 @@ proc findReachableNodes*(bootstrapNodes: seq[SignedPeerRecord]): seq[SignedPeerR ## Currently returns bootstrap nodes. In the future, any network participant ## confirmed reachable by AutoNAT could be included. bootstrapNodes + +# Hole punching logic below is adapted from libp2p's HPService +# (libp2p/services/hpservice.nim). HPService cannot be used directly because it +# depends on AutoNAT v1 and starts the relay immediately on NotReachable, +# bypassing the UPnP step. + +proc tryStartingDirectConn( + switch: Switch, peerId: PeerId +): Future[bool] {.async: (raises: [CancelledError]).} = + proc tryConnect( + address: MultiAddress + ): Future[bool] {.async: (raises: [DialFailedError, CancelledError]).} = + debug "Trying to create direct connection", peerId, address + await switch.connect(peerId, @[address], true, false) + debug "Direct connection created." + return true + + await sleepAsync(500.milliseconds) # wait for AddressBook to be populated + for address in switch.peerStore[AddressBook][peerId]: + try: + let isRelayedAddr = address.contains(multiCodec("p2p-circuit")) + if not isRelayedAddr.get(false) and address.isPublicMA(): + return await tryConnect(address) + except CatchableError as err: + debug "Failed to create direct connection.", description = err.msg + continue + return false + +proc closeRelayConn(relayedConn: Connection) {.async: (raises: [CancelledError]).} = + await sleepAsync(2000.milliseconds) # grace period before closing relayed connection + await relayedConn.close() + +proc holePunchIfRelayed*( + switch: Switch, peerId: PeerId +) {.async: (raises: [CancelledError]).} = + ## Attempts to establish a direct connection when a peer connected via relay. + ## First tries a direct TCP connect (if the peer's address is known and public), + ## then falls back to dcutr simultaneous-open hole punching. + ## Closes the relay connection once a direct path is established. + let connections = + switch.connManager.getConnections().getOrDefault(peerId).mapIt(it.connection) + if connections.anyIt(not isRelayed(it)): + return + let incomingRelays = connections.filterIt(it.transportDir == Direction.In) + if incomingRelays.len == 0: + return + + let relayedConn = incomingRelays[0] + + if await tryStartingDirectConn(switch, peerId): + await closeRelayConn(relayedConn) + return + + var natAddrs = switch.peerStore.getMostObservedProtosAndPorts() + if natAddrs.len == 0: + natAddrs = switch.peerInfo.listenAddrs.mapIt(switch.peerStore.guessDialableAddr(it)) + try: + await DcutrClient.new().startSync(switch, peerId, natAddrs) + await closeRelayConn(relayedConn) + except DcutrError as err: + debug "Hole punching failed during dcutr", description = err.msg + +proc setupHolePunching*(switch: Switch) = + try: + switch.mount(Dcutr.new(switch)) + except LPError as err: + error "Failed to mount Dcutr protocol", description = err.msg + + let handler = proc( + peerId: PeerId, event: PeerEvent + ) {.async: (raises: [CancelledError]).} = + await holePunchIfRelayed(switch, peerId) + switch.addPeerEventHandler(handler, PeerEventKind.Joined) diff --git a/storage/storage.nim b/storage/storage.nim index 450587cc..15165fcc 100644 --- a/storage/storage.nim +++ b/storage/storage.nim @@ -451,6 +451,8 @@ proc new*( ) ) + setupHolePunching(switch) + # REST server var restServer: RestServerRef = nil diff --git a/tests/storage/testnat.nim b/tests/storage/testnat.nim index 0808cf2d..e4ac92d3 100644 --- a/tests/storage/testnat.nim +++ b/tests/storage/testnat.nim @@ -3,6 +3,8 @@ import pkg/chronos import pkg/libp2p/[multiaddress, multihash, multicodec] import pkg/libp2p/protocols/connectivity/autonat/types import pkg/libp2p/protocols/connectivity/relay/client as relayClientModule +import pkg/libp2p/protocols/connectivity/dcutr/core as dcutrCore +import pkg/libp2p/multistream import pkg/libp2p/services/autorelayservice except setup import pkg/results @@ -48,8 +50,9 @@ asyncchecksuite "NAT - handleNatStatus": test "handleNatStatus announces mapped address when NotReachable and UPnP succeeds": let dialBack = MultiAddress.init("/ip4/1.2.3.4/tcp/8080").expect("valid") - let mapper = - MockNatPortMapper(mappedPorts: some((Port(9000), Port(9001), MappingProtocol.UPnP))) + let mapper = MockNatPortMapper( + mappedPorts: some((Port(9000), Port(9001), MappingProtocol.UPnP)) + ) await mapper.handleNatStatus( NotReachable, Opt.some(dialBack), discoveryPort, disc, sw, autoRelay @@ -105,3 +108,31 @@ asyncchecksuite "NAT - handleNatStatus": check not autoRelay.isRunning check disc.announceAddrs == @[dialBack] check not disc.protocol.clientMode + +asyncchecksuite "NAT - Hole punching": + test "setupHolePunching mounts the dcutr protocol on the switch": + let sw = newStandardSwitch() + setupHolePunching(sw) + check sw.ms.handlers.anyIt(dcutrCore.DcutrCodec in it.protos) + + test "holePunchIfRelayed returns early when the peer has no connections": + let sw1 = newStandardSwitch() + let sw2 = newStandardSwitch() + await allFutures(sw1.start(), sw2.start()) + + await holePunchIfRelayed(sw1, sw2.peerInfo.peerId) + + await allFutures(sw1.stop(), sw2.stop()) + + test "holePunchIfRelayed returns early when a direct connection already exists": + let sw1 = newStandardSwitch() + let sw2 = newStandardSwitch() + await allFutures(sw1.start(), sw2.start()) + + await sw1.connect(sw2.peerInfo.peerId, sw2.peerInfo.addrs) + check sw1.isConnected(sw2.peerInfo.peerId) + + await holePunchIfRelayed(sw1, sw2.peerInfo.peerId) + + check sw1.isConnected(sw2.peerInfo.peerId) + await allFutures(sw1.stop(), sw2.stop())