diff --git a/docker-compose.hole-punch.yml b/docker-compose.hole-punch.yml new file mode 100644 index 0000000..29eccdf --- /dev/null +++ b/docker-compose.hole-punch.yml @@ -0,0 +1,138 @@ +# Local reproduction of the libp2p/unified-testing hole-punch topology: +# two peers on isolated LANs, two Linux NAT routers, one WAN relay and Redis +# coordination on a separate network. +# +# Usage: docker compose -f docker-compose.hole-punch.yml up --build \ +# --exit-code-from hs-hp-dialer +x-hs-env: &hs-env + REDIS_ADDR: "redis:6379" + TEST_KEY: "cafe0131" + TRANSPORT: tcp + SECURE_CHANNEL: noise + MUXER: yamux + +x-router: &router + build: + context: . + dockerfile: interop/hole-punch/router/Dockerfile + cap_add: [NET_ADMIN] + sysctls: + net.ipv4.ip_forward: 1 + net.ipv4.conf.all.forwarding: 1 + net.ipv4.conf.default.forwarding: 1 + net.ipv4.conf.all.rp_filter: 0 + net.ipv4.conf.default.rp_filter: 0 + +networks: + wan: + ipam: + config: [{subnet: 10.131.49.64/27}] + lan-dialer: + ipam: + config: [{subnet: 10.131.49.96/27}] + lan-listener: + ipam: + config: [{subnet: 10.131.49.128/27}] + coordination: {} + +services: + redis: + image: redis:7-alpine + command: ["redis-server", "--save", "", "--appendonly", "no"] + networks: + - coordination + healthcheck: + test: ["CMD", "redis-cli", "ping"] + interval: 1s + timeout: 3s + retries: 30 + + hs-hp-relay: + build: + context: . + dockerfile: interop/hole-punch/Dockerfile + depends_on: + redis: + condition: service_healthy + cap_add: [NET_ADMIN] + networks: + wan: + ipv4_address: 10.131.49.68 + interface_name: wan0 + coordination: + interface_name: redis0 + environment: + <<: *hs-env + RELAY_IP: "10.131.49.68" + DIALER_LAN_SUBNET: "10.131.49.96/27" + DIALER_ROUTER_IP: "10.131.49.66" + LISTENER_LAN_SUBNET: "10.131.49.128/27" + LISTENER_ROUTER_IP: "10.131.49.67" + + dialer-router: + <<: *router + networks: + wan: + ipv4_address: 10.131.49.66 + interface_name: wan0 + lan-dialer: + ipv4_address: 10.131.49.98 + interface_name: lan0 + environment: + WAN_IP: "10.131.49.66" + WAN_SUBNET: "10.131.49.64/27" + LAN_IP: "10.131.49.98" + LAN_SUBNET: "10.131.49.96/27" + + listener-router: + <<: *router + networks: + wan: + ipv4_address: 10.131.49.67 + interface_name: wan0 + lan-listener: + ipv4_address: 10.131.49.130 + interface_name: lan0 + environment: + WAN_IP: "10.131.49.67" + WAN_SUBNET: "10.131.49.64/27" + LAN_IP: "10.131.49.130" + LAN_SUBNET: "10.131.49.128/27" + + hs-hp-listener: + build: + context: . + dockerfile: interop/hole-punch/Dockerfile + depends_on: [redis, hs-hp-relay, listener-router] + cap_add: [NET_ADMIN] + networks: + lan-listener: + ipv4_address: 10.131.49.131 + interface_name: lan0 + coordination: + interface_name: redis0 + environment: + <<: *hs-env + IS_DIALER: "false" + LISTENER_IP: "10.131.49.131" + WAN_SUBNET: "10.131.49.64/27" + WAN_ROUTER_IP: "10.131.49.130" + + hs-hp-dialer: + build: + context: . + dockerfile: interop/hole-punch/Dockerfile + depends_on: [redis, hs-hp-relay, dialer-router, hs-hp-listener] + cap_add: [NET_ADMIN] + networks: + lan-dialer: + ipv4_address: 10.131.49.99 + interface_name: lan0 + coordination: + interface_name: redis0 + environment: + <<: *hs-env + IS_DIALER: "true" + DIALER_IP: "10.131.49.99" + WAN_SUBNET: "10.131.49.64/27" + WAN_ROUTER_IP: "10.131.49.98" diff --git a/interop/Makefile b/interop/Makefile index fe0b2f9..baea769 100644 --- a/interop/Makefile +++ b/interop/Makefile @@ -7,7 +7,8 @@ IMAGE_NAME := libp2p-hs-interop .PHONY: build image self-test cross-test-go-listener cross-test-hs-listener \ cross-gossipsub-rust-listener cross-gossipsub-hs-listener cross-gossipsub kad-dht \ - perf-self-test perf-cross-nim-listener perf-cross-hs-listener perf-cross nim-libp2p + perf-self-test perf-cross-nim-listener perf-cross-hs-listener perf-cross nim-libp2p \ + hole-punch-self-test # nim-libp2p commit pinned by unified-testing perf/images.yaml (nim-v1.15). NIM_LIBP2P_COMMIT := 1bdf2f67971529e8bee01252230bdb00ab785ef7 @@ -63,3 +64,9 @@ perf-cross-hs-listener: nim-libp2p # Perf cross-test: both directions perf-cross: perf-cross-nim-listener perf-cross-hs-listener + +# Hole-punch self-test: hs peers behind separate NAT routers with a WAN relay +hole-punch-self-test: + cd .. && docker compose -f docker-compose.hole-punch.yml up --build \ + --exit-code-from hs-hp-dialer redis dialer-router listener-router \ + hs-hp-relay hs-hp-listener hs-hp-dialer diff --git a/interop/hole-punch/Dockerfile b/interop/hole-punch/Dockerfile new file mode 100644 index 0000000..7cdac0b --- /dev/null +++ b/interop/hole-punch/Dockerfile @@ -0,0 +1,40 @@ +# Multi-stage build for the libp2p-hs hole-punch test daemon. +# Used by the libp2p/unified-testing hole-punch framework; build context +# is the repository root. + +# Stage 1: Build with GHC 9.10 +FROM haskell:9.10-slim-bookworm AS builder + +WORKDIR /app + +COPY libp2p-hs.cabal cabal.project ./ + +# Fix ppad-sha256 ARM SHA2 intrinsic compilation on Docker (GCC 12). +RUN if [ "$(uname -m)" = "aarch64" ]; then \ + echo 'package ppad-sha256' >> cabal.project && \ + echo ' ghc-options: -optc-march=armv8-a+crypto' >> cabal.project; \ + fi + +RUN cabal update && cabal build --only-dependencies lib:libp2p-hs + +COPY src/ src/ +COPY interop/ interop/ +RUN cabal build libp2p-hole-punch \ + && cp "$(cabal list-bin libp2p-hole-punch)" /app/libp2p-hole-punch \ + && strip /app/libp2p-hole-punch + +# Stage 2: Minimal runtime image +FROM debian:bookworm-slim + +RUN apt-get update \ + && apt-get install -y --no-install-recommends \ + libgmp10 \ + ca-certificates \ + iproute2 \ + && rm -rf /var/lib/apt/lists/* + +COPY --from=builder /app/libp2p-hole-punch /usr/local/bin/libp2p-hole-punch +COPY interop/hole-punch/entrypoint.sh /usr/local/bin/hole-punch-entrypoint +RUN chmod +x /usr/local/bin/hole-punch-entrypoint + +ENTRYPOINT ["/usr/local/bin/hole-punch-entrypoint"] diff --git a/interop/hole-punch/Main.hs b/interop/hole-punch/Main.hs new file mode 100644 index 0000000..e5c1772 --- /dev/null +++ b/interop/hole-punch/Main.hs @@ -0,0 +1,343 @@ +-- | Hole-punch / DCUtR test daemon for libp2p/unified-testing. +-- +-- Implements docs/write-a-hole-punch-test-app.md: Redis coordination +-- namespaced by TEST_KEY, three roles (relay / dialer / listener), +-- YAML results on stdout (dialer only), logging on stderr. +module Main (main) where + +import Control.Applicative ((<|>)) +import Control.Concurrent (threadDelay) +import Control.Concurrent.STM (atomically, readTVar) +import Control.Monad (filterM, forever) +import qualified Data.ByteString.Char8 as BS8 +import Data.List (find) +import Data.Maybe (fromMaybe) +import qualified Data.Text as T +import qualified Data.Text.Encoding as TE +import Data.Time.Clock (UTCTime, diffUTCTime, getCurrentTime) +import qualified Database.Redis as Redis +import LibP2P + ( Connection (..) + , Multiaddr (..) + , PeerId + , PingResult (..) + , Protocol (..) + , Switch + , addTransport + , defaultConnectionGater + , defaultNATConfig + , dial + , fromPublicKey + , fromText + , generateKeyPair + , newSwitch + , newTCPTransport + , parsePeerId + , peerIdBytes + , registerIdentifyHandlers + , registerNATHandlers + , registerPingHandler + , sendPing + , switchClose + , switchListen + , toBase58 + , toText + ) +import LibP2P.Crypto.Key (publicKey) +import LibP2P.Multiaddr (encapsulate, isRelayedAddr) +import LibP2P.NAT.Relay.Transport (CircuitAddr (..), parseCircuitAddr) +import LibP2P.Protocol.Identify (identifyPeer) +import LibP2P.Switch.ConnPool (lookupAllConns) +import LibP2P.Switch.Dial (DialOpts (..), defaultDialOpts, dialWith) +import LibP2P.Switch.Types (ConnState (..), swConnPool, swLocalPeerId) +import Network.Socket + ( AddrInfo (..) + , SockAddr (..) + , defaultHints + , getAddrInfo + , hostAddressToTuple + ) +import qualified Network.Socket as Socket +import System.Environment (lookupEnv) +import System.Exit (exitFailure, exitSuccess) +import System.IO (hFlush, hPutStrLn, stderr, stdout) +import System.Timeout (timeout) +import Text.Printf (printf) +import Text.Read (readMaybe) + +coordinationTimeoutSeconds :: Int +coordinationTimeoutSeconds = 150 + +dialerTimeoutSeconds :: Int +dialerTimeoutSeconds = 170 + +main :: IO () +main = do + redisAddr <- fromMaybe "hole-punch-redis:6379" <$> lookupEnv "REDIS_ADDR" + testKey <- getEnvRequired "TEST_KEY" + transport <- getEnvRequired "TRANSPORT" + security <- lookupEnv "SECURE_CHANNEL" + muxer <- lookupEnv "MUXER" + case validateProtocols transport security muxer of + Left err -> die err + Right () -> pure () + sw <- newNode + redisConn <- connectRedis redisAddr + isRelay <- lookupEnv "IS_RELAY" + isDialer <- lookupEnv "IS_DIALER" + case (isRelay, isDialer) of + (Just "true", _) -> runRelay sw redisConn testKey + (_, Just "true") -> runDialerBounded sw redisConn testKey + (_, Just "false") -> runListener sw redisConn testKey + -- The executable harness uses a dedicated relay image and currently + -- omits IS_RELAY. Peer containers always receive IS_DIALER. + (Nothing, Nothing) -> runRelay sw redisConn testKey + other -> dieWith sw ("Invalid role environment: " ++ show other) + +newNode :: IO Switch +newNode = do + ekp <- generateKeyPair + kp <- either die pure ekp + let pid = fromPublicKey (publicKey kp) + logInfo $ "PeerId: " ++ T.unpack (toBase58 pid) + sw <- newSwitch pid kp + tcp <- newTCPTransport + addTransport sw tcp + registerIdentifyHandlers sw + registerPingHandler sw + _ <- registerNATHandlers sw defaultNATConfig + pure sw + +runRelay :: Switch -> Redis.Connection -> String -> IO () +runRelay sw redisConn testKey = do + ip <- fromMaybe "0.0.0.0" <$> lookupEnv "RELAY_IP" + addrText <- listenTcp sw ip + redisSet redisConn (redisKey testKey "relay_multiaddr") addrText + logInfo $ "Relay listening on " ++ T.unpack addrText + forever $ threadDelay 3600000000 + +runListener :: Switch -> Redis.Connection -> String -> IO () +runListener sw redisConn testKey = do + _ <- listenTcp sw =<< nodeIp "LISTENER_IP" + relayMA <- waitRelayAddr redisConn testKey sw + reserveOnRelay sw relayMA + redisSet redisConn (redisKey testKey "listener_peer_id") (toBase58 (swLocalPeerId sw)) + logInfo $ "Published listener peer id " ++ T.unpack (toBase58 (swLocalPeerId sw)) + forever $ threadDelay 3600000000 + +runDialerBounded :: Switch -> Redis.Connection -> String -> IO () +runDialerBounded sw redisConn testKey = do + result <- timeout (dialerTimeoutSeconds * 1000000) (runDialer sw redisConn testKey) + case result of + Nothing -> dieWith sw "Hole-punch test exceeded the 170 second deadline" + Just () -> pure () + +runDialer :: Switch -> Redis.Connection -> String -> IO () +runDialer sw redisConn testKey = do + _ <- listenTcp sw =<< nodeIp "DIALER_IP" + relayMA <- waitRelayAddr redisConn testKey sw + reserveOnRelay sw relayMA + listenerId <- waitListenerId redisConn testKey sw + logInfo $ "Dialing listener via relay: " ++ T.unpack (toBase58 listenerId) + t0 <- getCurrentTime + let circuitAddr = + encapsulate relayMA (Multiaddr [P2PCircuit, P2P (peerIdBytes listenerId)]) + _ <- dial sw listenerId [circuitAddr] + >>= either (\err -> dieWith sw ("Circuit dial failed: " ++ show err)) pure + direct <- waitDirectConn sw listenerId (coordinationTimeoutSeconds * 5) + case direct of + Nothing -> dieWith sw "DCUtR failed: no direct connection within timeout" + Just conn -> finishDial sw conn t0 + +finishDial :: Switch -> Connection -> UTCTime -> IO () +finishDial sw conn t0 = do + pingResult <- sendPing sw conn + case pingResult of + Left err -> dieWith sw ("Ping over direct connection failed: " ++ show err) + Right result -> do + completedAt <- getCurrentTime + let handshakePlusRTT = toMilliseconds (diffUTCTime completedAt t0) + pingRTTMillis = toMilliseconds (pingRTT result) + logInfo "Direct connection established and verified via DCUtR" + printf "latency:\n handshake_plus_one_rtt: %.2f\n ping_rtt: %.2f\n unit: ms\n" + handshakePlusRTT pingRTTMillis + hFlush stdout + switchClose sw + exitSuccess + +-- | Convert a duration in seconds to milliseconds for the harness schema. +toMilliseconds :: Real a => a -> Double +toMilliseconds value = realToFrac value * 1000 + +listenTcp :: Switch -> String -> IO T.Text +listenTcp sw ip = do + bindAddr <- either (\err -> dieWith sw ("Invalid bind address: " ++ err)) pure + (fromText (T.pack ("/ip4/" ++ ip ++ "/tcp/0"))) + addrs <- switchListen sw defaultConnectionGater [bindAddr] + case addrs of + [] -> dieWith sw "switchListen returned no TCP addresses" + (listenAddr : _) -> do + actual <- resolveListenAddr listenAddr ip + let full = encapsulate actual (Multiaddr [P2P (peerIdBytes (swLocalPeerId sw))]) + logInfo $ "Listening on " ++ T.unpack (toText full) + pure (toText full) + +reserveOnRelay :: Switch -> Multiaddr -> IO () +reserveOnRelay sw relayMA = do + let circuitListen = encapsulate relayMA (Multiaddr [P2PCircuit]) + relayId <- either (dieWith sw . ("Invalid relay address: " ++)) (pure . caRelayId) + (parseCircuitAddr circuitListen) + -- Establish the relay mapping from the TCP listen port. Identify's + -- observed address then names the same mapping used by simultaneous open. + let opts = defaultDialOpts { doForceDirect = True } + relayConn <- dialWith sw opts relayId [relayMA] + >>= either (dieWith sw . ("Relay dial failed: " ++) . show) pure + identifyPeer sw relayConn + >>= either (dieWith sw . ("Relay Identify failed: " ++)) pure + _ <- switchListen sw defaultConnectionGater [circuitListen] + logInfo $ "Reserved on relay " ++ T.unpack (toText relayMA) + +waitRelayAddr :: Redis.Connection -> String -> Switch -> IO Multiaddr +waitRelayAddr redisConn testKey sw = do + raw <- pollRedis redisConn (redisKey testKey "relay_multiaddr") + >>= maybe (dieWith sw "Timed out waiting for relay multiaddr") pure + either (\err -> dieWith sw ("Bad relay multiaddr: " ++ err)) pure (fromText (TE.decodeUtf8 raw)) + +waitListenerId :: Redis.Connection -> String -> Switch -> IO PeerId +waitListenerId redisConn testKey sw = do + raw <- pollRedis redisConn (redisKey testKey "listener_peer_id") + >>= maybe (dieWith sw "Timed out waiting for listener peer id") pure + either (\err -> dieWith sw ("Failed to parse listener peer id: " ++ err)) pure + (parsePeerId (TE.decodeUtf8 raw)) + +waitDirectConn :: Switch -> PeerId -> Int -> IO (Maybe Connection) +waitDirectConn sw pid attempts = go attempts + where + go 0 = pure Nothing + go n = do + conns <- atomically $ lookupAllConns (swConnPool sw) pid + openDirect <- filterM isOpenDirect conns + case openDirect of + (c : _) -> pure (Just c) + [] -> threadDelay 200000 >> go (n - 1) + +isOpenDirect :: Connection -> IO Bool +isOpenDirect c = do + st <- atomically $ readTVar (connState c) + pure (st == ConnOpen && not (isRelayedAddr (connRemoteAddr c))) + +nodeIp :: String -> IO String +nodeIp roleVariable = do + roleIp <- lookupEnv roleVariable + fallback <- lookupEnv "PEER_IP" + pure (fromMaybe "0.0.0.0" (roleIp <|> fallback)) + +redisKey :: String -> String -> BS8.ByteString +redisKey testKey suffix = BS8.pack (testKey ++ "_" ++ suffix) + +connectRedis :: String -> IO Redis.Connection +connectRedis redisAddr = do + (host, port) <- either die pure (parseHostPort redisAddr) + Redis.checkedConnect Redis.defaultConnectInfo + { Redis.connectHost = host + , Redis.connectPort = Redis.PortNumber (fromIntegral port) + } + +redisSet :: Redis.Connection -> BS8.ByteString -> T.Text -> IO () +redisSet conn key value = do + result <- Redis.runRedis conn $ Redis.set key (TE.encodeUtf8 value) + case result of + Left err -> die ("Redis SET failed: " ++ show err) + Right _ -> pure () + +pollRedis :: Redis.Connection -> BS8.ByteString -> IO (Maybe BS8.ByteString) +pollRedis conn key = go (coordinationTimeoutSeconds * 2) + where + go 0 = pure Nothing + go n = do + result <- Redis.runRedis conn $ Redis.get key + case result of + Right (Just value) -> pure (Just value) + _ -> threadDelay 500000 >> go (n - 1) + +validateProtocols :: String -> Maybe String -> Maybe String -> Either String () +validateProtocols transport security muxer = do + case transport of + "tcp" -> pure () + other -> Left $ "transport " ++ other ++ " not supported (only tcp)" + case security of + Just "noise" -> pure () + Just other -> Left $ "secure channel " ++ other ++ " not supported (only noise)" + Nothing -> Left "SECURE_CHANNEL not set (required for tcp)" + case muxer of + Just "yamux" -> pure () + Just other -> Left $ "muxer " ++ other ++ " not supported (only yamux)" + Nothing -> Left "MUXER not set (required for tcp)" + +parseHostPort :: String -> Either String (String, Int) +parseHostPort value = case break (== ':') value of + (host, []) + | null host -> Left "REDIS_ADDR host is empty" + | otherwise -> Right (host, 6379) + (host, ':' : portText) + | null host -> Left "REDIS_ADDR host is empty" + | ':' `elem` portText -> Left "REDIS_ADDR must use host:port" + | otherwise -> case readMaybe portText of + Just port | port >= 1 && port <= 65535 -> Right (host, port) + _ -> Left "REDIS_ADDR port must be an integer from 1 to 65535" + _ -> Left "Invalid REDIS_ADDR" + +resolveListenAddr :: Multiaddr -> String -> IO Multiaddr +resolveListenAddr addr ip + | ip == "0.0.0.0" = do + actualIP <- discoverContainerIP + case addr of + Multiaddr (IP4 _ : rest) -> + case fromText (T.pack ("/ip4/" ++ actualIP)) of + Right (Multiaddr [IP4 w]) -> pure $ Multiaddr (IP4 w : rest) + _ -> pure addr + _ -> pure addr + | otherwise = pure addr + +discoverContainerIP :: IO String +discoverContainerIP = do + mHostname <- lookupEnv "HOSTNAME" + case mHostname of + Nothing -> pure "0.0.0.0" + Just hostname -> do + addrs <- getAddrInfo (Just defaultHints) (Just hostname) Nothing :: IO [AddrInfo] + case find isNonLoopbackIPv4 addrs of + Just ai -> pure $ sockAddrToIP (Socket.addrAddress ai) + Nothing -> pure "0.0.0.0" + +sockAddrToIP :: SockAddr -> String +sockAddrToIP (SockAddrInet _ hostAddr) = + let (a, b, c, d) = hostAddressToTuple hostAddr + in show a ++ "." ++ show b ++ "." ++ show c ++ "." ++ show d +sockAddrToIP _ = "0.0.0.0" + +isNonLoopbackIPv4 :: AddrInfo -> Bool +isNonLoopbackIPv4 ai = case Socket.addrAddress ai of + SockAddrInet _ hostAddr -> + let (a, _, _, _) = hostAddressToTuple hostAddr + in a /= 127 + _ -> False + +getEnvRequired :: String -> IO String +getEnvRequired name = do + val <- lookupEnv name + case val of + Just v -> pure v + Nothing -> die ("Missing required environment variable: " ++ name) + +dieWith :: Switch -> String -> IO a +dieWith sw msg = do + hPutStrLn stderr msg + switchClose sw + exitFailure + +die :: String -> IO a +die msg = hPutStrLn stderr msg >> exitFailure + +logInfo :: String -> IO () +logInfo msg = hPutStrLn stderr msg >> hFlush stderr diff --git a/interop/hole-punch/entrypoint.sh b/interop/hole-punch/entrypoint.sh new file mode 100644 index 0000000..aadb780 --- /dev/null +++ b/interop/hole-punch/entrypoint.sh @@ -0,0 +1,33 @@ +#!/bin/sh +set -eu + +add_route() { + subnet="$1" + gateway="$2" + interface="$3" + label="$4" + + if [ -z "$subnet" ] || [ -z "$gateway" ]; then + return + fi + + echo "Setting route to ${label} ${subnet} via ${gateway}" >&2 + ip route replace "$subnet" via "$gateway" dev "$interface" +} + +if [ -n "${IS_DIALER+x}" ]; then + add_route "${WAN_SUBNET:-}" "${WAN_ROUTER_IP:-}" lan0 "WAN subnet" +else + add_route \ + "${DIALER_LAN_SUBNET:-}" \ + "${DIALER_ROUTER_IP:-}" \ + wan0 \ + "dialer LAN" + add_route \ + "${LISTENER_LAN_SUBNET:-}" \ + "${LISTENER_ROUTER_IP:-}" \ + wan0 \ + "listener LAN" +fi + +exec /usr/local/bin/libp2p-hole-punch "$@" diff --git a/interop/hole-punch/router/Dockerfile b/interop/hole-punch/router/Dockerfile new file mode 100644 index 0000000..5fb81f4 --- /dev/null +++ b/interop/hole-punch/router/Dockerfile @@ -0,0 +1,10 @@ +FROM alpine:3.19 + +RUN apk add --no-cache \ + iproute2 \ + iptables + +COPY interop/hole-punch/router/run.sh /usr/local/bin/run-hole-punch-router +RUN chmod +x /usr/local/bin/run-hole-punch-router + +ENTRYPOINT ["/usr/local/bin/run-hole-punch-router"] diff --git a/interop/hole-punch/router/run.sh b/interop/hole-punch/router/run.sh new file mode 100644 index 0000000..816b4dd --- /dev/null +++ b/interop/hole-punch/router/run.sh @@ -0,0 +1,33 @@ +#!/bin/sh +set -eu + +require_env() { + name="$1" + eval "value=\${$name:-}" + if [ -z "$value" ]; then + echo "Missing required environment variable: $name" >&2 + exit 1 + fi +} + +require_env WAN_IP +require_env WAN_SUBNET +require_env LAN_IP +require_env LAN_SUBNET + +iptables -t nat -F +iptables -t filter -F +iptables -P FORWARD DROP +iptables -P INPUT ACCEPT +iptables -P OUTPUT ACCEPT + +# Do not reject the first SYN before the peer creates its matching mapping. +iptables -A INPUT -i wan0 -m conntrack --ctstate ESTABLISHED,RELATED -j ACCEPT +iptables -A INPUT -i wan0 -j DROP +iptables -t nat -A POSTROUTING -s "$LAN_SUBNET" -o wan0 -j MASQUERADE +iptables -A FORWARD -m conntrack --ctstate ESTABLISHED,RELATED -j ACCEPT +iptables -A FORWARD -s "$LAN_SUBNET" -i lan0 -o wan0 -j ACCEPT +iptables -A FORWARD -d "$LAN_SUBNET" -i wan0 -o lan0 -j ACCEPT + +echo "NAT router ready: ${LAN_SUBNET} via ${WAN_IP}" >&2 +exec tail -f /dev/null diff --git a/libp2p-hs.cabal b/libp2p-hs.cabal index b44283e..ab57bc2 100644 --- a/libp2p-hs.cabal +++ b/libp2p-hs.cabal @@ -147,6 +147,21 @@ executable libp2p-perf libp2p-hs ghc-options: -threaded -rtsopts +executable libp2p-hole-punch + import: warnings, lang + hs-source-dirs: interop/hole-punch + main-is: Main.hs + build-depends: + base >= 4.18 && < 5, + bytestring >= 0.10 && < 0.13, + text >= 1.2 && < 2.2, + time >= 1.9 && < 2, + network >= 3.1 && < 3.3, + hedis >= 0.15 && < 0.16, + stm >= 2.4 && < 2.6, + libp2p-hs + ghc-options: -threaded -rtsopts + executable libp2p-kad-dht-node import: warnings, lang hs-source-dirs: interop/kad-dht-node diff --git a/src/LibP2P/NAT.hs b/src/LibP2P/NAT.hs index 0e883f9..29214ce 100644 --- a/src/LibP2P/NAT.hs +++ b/src/LibP2P/NAT.hs @@ -23,6 +23,7 @@ module LibP2P.NAT , registerDCUtRUpgrade , upgradeRelayedConnection , holePunchTargets + , dcutrOwnAddrs , DCUtRUpgradeConfig (..) , defaultDCUtRUpgradeConfig -- * Circuit client @@ -35,6 +36,7 @@ import Control.Concurrent (threadDelay) import Control.Concurrent.Async (async) import Control.Concurrent.STM (atomically, modifyTVar', readTVar) import Control.Monad (filterM, unless, void) +import Data.List (nub) import Data.Maybe (fromMaybe) import System.Timeout (timeout) import qualified Data.Map.Strict as Map @@ -73,12 +75,14 @@ import LibP2P.NAT.Relay.Message , writeHopMessage ) import LibP2P.NAT.Relay.Transport - ( CircuitState + ( CircuitAddr (..) + , CircuitState , ReservationRefreshConfig (..) , acceptStopStream , circuitTransport , defaultReservationRefreshConfig , newCircuitState + , parseCircuitAddr ) import LibP2P.Switch (addTransport, selectTransport, setStreamHandler) import LibP2P.Switch.ConnPool (lookupAllConns, lookupConn) @@ -301,7 +305,7 @@ initiateOverRelay sw config relayConn = do closeQuietly stream pure (DCUtRFailed "remote does not support /libp2p/dcutr") Accepted _ -> do - ownAddrs <- dialableListenAddrs sw + ownAddrs <- dcutrOwnAddrs sw (relayObserver relayConn) let dcConfig = DCUtRConfig { dcMaxAttempts = ducMaxAttempts config , dcDialer = \addr -> @@ -341,9 +345,36 @@ holePunchDial sw config asClient peerId addrs = do Right (Just (Left err)) -> Left (show err) Right (Just (Right _conn)) -> Right () --- | Our own listen addresses that a peer could hole punch to. -dialableListenAddrs :: Switch -> IO [Multiaddr] -dialableListenAddrs sw = filter (not . isRelayedAddr) <$> switchListenAddrs sw +-- | Addresses we put in DCUtR CONNECT: the address reported by the relay +-- that carries this connection, followed by every non-relayed listen +-- address. Scoping the observation to that relay avoids advertising stale +-- mappings learned from unrelated peers. Private listen addresses remain +-- useful to peers on the same LAN and as deterministic test fallbacks. +dcutrOwnAddrs :: Switch -> Maybe PeerId -> IO [Multiaddr] +dcutrOwnAddrs sw observer = do + observed <- observedAddrs sw observer + listen <- filter (not . isRelayedAddr) <$> switchListenAddrs sw + pure (nub (observed ++ listen)) + +-- | How the relevant relay observed us during Identify. +observedAddrs :: Switch -> Maybe PeerId -> IO [Multiaddr] +observedAddrs sw observer = do + store <- atomically $ readTVar (swPeerStore sw) + let infos = case observer of + Just peerId -> maybe [] pure (Map.lookup peerId store) + Nothing -> [] + pure + [ addr + | info <- infos + , Just raw <- [idObservedAddr info] + , Right addr <- [fromBytes raw] + , not (isRelayedAddr addr) + ] + +-- | The relay encoded in a relayed connection's remote multiaddr. +relayObserver :: Connection -> Maybe PeerId +relayObserver conn = + either (const Nothing) (Just . caRelayId) (parseCircuitAddr (connRemoteAddr conn)) -- | Close the relay connection after the grace period, provided a direct -- connection to the peer is still up. @@ -500,7 +531,7 @@ registerRelayStopHandler sw circuitState = registerDCUtRHandler :: Switch -> DCUtRUpgradeConfig -> IO () registerDCUtRHandler sw upgradeConfig = setStreamHandler sw dcutrProtocolId $ \conn stream -> do - addrs <- dialableListenAddrs sw + addrs <- dcutrOwnAddrs sw (relayObserver conn) let config = DCUtRConfig { dcMaxAttempts = ducMaxAttempts upgradeConfig -- We are peer A: the spec makes us the client of the diff --git a/src/LibP2P/NAT/Relay/Transport.hs b/src/LibP2P/NAT/Relay/Transport.hs index 3b8c898..235286b 100644 --- a/src/LibP2P/NAT/Relay/Transport.hs +++ b/src/LibP2P/NAT/Relay/Transport.hs @@ -130,9 +130,10 @@ defaultReservationRefreshConfig = ReservationRefreshConfig -- 'LibP2P.Switch.newSwitch' with 'LibP2P.Switch.addTransport'. circuitTransport :: Switch -> CircuitState -> ReservationRefreshConfig -> Transport circuitTransport sw st refreshCfg = Transport - { transportDial = dialCircuit sw - , transportListen = listenCircuit sw st refreshCfg - , transportCanDial = either (const False) (const True) . parseCircuitAddr + { transportDial = dialCircuit sw + , transportDialFrom = \_ -> dialCircuit sw + , transportListen = listenCircuit sw st refreshCfg + , transportCanDial = either (const False) (const True) . parseCircuitAddr } -- Address handling diff --git a/src/LibP2P/Switch/Dial.hs b/src/LibP2P/Switch/Dial.hs index 646a9eb..c933593 100644 --- a/src/LibP2P/Switch/Dial.hs +++ b/src/LibP2P/Switch/Dial.hs @@ -28,7 +28,7 @@ module LibP2P.Switch.Dial ) where import Control.Concurrent (threadDelay) -import Control.Concurrent.Async (Async, async, cancel, waitAnyCatch) +import Control.Concurrent.Async (Async, async, cancel, waitAnyCatch, waitCatch) import Control.Concurrent.STM ( STM , TMVar @@ -42,16 +42,17 @@ import Control.Concurrent.STM , writeTChan , writeTVar ) -import Control.Exception (SomeException, finally, onException) -import Control.Monad (forM, when) +import Control.Exception (SomeException, bracketOnError, catch, finally, onException) +import Control.Monad (forM, forM_, when) import Data.List (find) import qualified Data.Map.Strict as Map import Data.Time.Clock (NominalDiffTime, addUTCTime, getCurrentTime) import LibP2P.Crypto.PeerId (PeerId) -import LibP2P.Multiaddr (Multiaddr) +import LibP2P.Multiaddr (Multiaddr (..), isRelayedAddr) +import LibP2P.Multiaddr.Protocol (Protocol (..)) import LibP2P.Switch.ConnPool (addConn, lookupConn) import LibP2P.Switch.Connection (closeConnection) -import LibP2P.Switch.Listen (streamAcceptLoop) +import LibP2P.Switch.Listen (streamAcceptLoop, switchListenAddrs) import LibP2P.Switch.ResourceManager (Direction (..), releaseConnection, reserveConnection) import LibP2P.Switch.Types ( BackoffEntry (..) @@ -62,7 +63,7 @@ import LibP2P.Switch.Types , SwitchEvent (..) ) import LibP2P.Switch.Upgrade (upgradeAs) -import LibP2P.Transport (Transport (..)) +import LibP2P.Transport (RawConnection (..), Transport (..)) -- | Initial backoff duration after first failure: 5 seconds. initialBackoffSeconds :: NominalDiffTime @@ -258,7 +259,8 @@ establishAndRegister sw opts remotePeerId addrs = do case resCheck of Left resErr -> pure (Left (DialResourceLimit resErr)) Right () -> do - result <- dialNewInner sw dir addrs + result <- dialNewInner sw opts dir addrs + `onException` atomically (releaseConnection (swResourceMgr sw) remotePeerId dir) -- Verify remote PeerId matches expected target let verified = case result of Right conn @@ -291,9 +293,9 @@ establishAndRegister sw opts remotePeerId addrs = do pure verified -- | Inner dial logic: transport selection and staggered parallel dial. -dialNewInner :: Switch -> Direction -> [Multiaddr] -> IO (Either DialError Connection) -dialNewInner _sw _dir [] = pure (Left DialNoAddresses) -dialNewInner sw dir addrs = do +dialNewInner :: Switch -> DialOpts -> Direction -> [Multiaddr] -> IO (Either DialError Connection) +dialNewInner _sw _opts _dir [] = pure (Left DialNoAddresses) +dialNewInner sw opts dir addrs = do transports <- atomically $ readTVar (swTransports sw) -- Find a transport for each address let dialable = filterMap (\addr -> @@ -302,7 +304,7 @@ dialNewInner sw dir addrs = do Nothing -> Nothing) addrs case dialable of [] -> pure (Left (DialNoTransport (Prelude.head addrs))) - pairs -> staggeredDial sw dir pairs + pairs -> staggeredDial sw opts dir pairs -- | Filter and map a list, keeping only Just results. filterMap :: (a -> Maybe b) -> [a] -> [b] @@ -315,16 +317,36 @@ filterMap f (x:xs) = case f x of -- -- Addresses are tried with 250ms delay between each attempt. -- The first successful connection wins; remaining attempts are cancelled. -staggeredDial :: Switch -> Direction -> [(Multiaddr, Transport)] -> IO (Either DialError Connection) -staggeredDial sw dir pairs = do - -- Spawn workers with staggered delays: 0ms, 250ms, 500ms, ... - workers <- forM (zip [0 :: Int ..] pairs) $ \(i, (addr, transport)) -> - async $ do - when (i > 0) $ threadDelay (i * staggerDelayUs) - rawConn <- transportDial transport addr - upgradeAs dir (swIdentityKey sw) rawConn - -- Wait for first success or collect all failures - collectResults workers [] +staggeredDial + :: Switch -> DialOpts -> Direction -> [(Multiaddr, Transport)] + -> IO (Either DialError Connection) +staggeredDial sw opts dir pairs = + bracketOnError spawnWorkers cancelAndCloseWorkers (`collectResults` []) + where + -- Spawn workers with staggered delays: 0ms, 250ms, 500ms, ... + spawnWorkers = forM (zip [0 :: Int ..] pairs) $ \(i, (addr, transport)) -> + async $ do + when (i > 0) $ threadDelay (i * staggerDelayUs) + localBind <- localBindFor sw (doForceDirect opts) addr + rawConn <- transportDialFrom transport localBind addr + upgradeAs dir (swIdentityKey sw) rawConn + `onException` rcClose rawConn + +-- | Hole-punch dials bind the outgoing socket to a same-family listen +-- address so the SYN shares the listen port. Ordinary dials leave the +-- source port ephemeral. +localBindFor :: Switch -> Bool -> Multiaddr -> IO (Maybe Multiaddr) +localBindFor _ False _ = pure Nothing +localBindFor sw True remote = do + addrs <- switchListenAddrs sw + pure $ find (sameIpFamily remote) (filter (not . isRelayedAddr) addrs) + +sameIpFamily :: Multiaddr -> Multiaddr -> Bool +sameIpFamily a b = ipKind a == ipKind b && ipKind a /= Nothing + where + ipKind (Multiaddr (IP4 _ : _)) = Just (0 :: Int) + ipKind (Multiaddr (IP6 _ : _)) = Just 1 + ipKind _ = Nothing -- | Wait for the first successful async result, cancelling the rest. -- If all fail, return DialAllFailed with all error messages. @@ -335,7 +357,19 @@ collectResults workers errs = do let remaining = filter (/= completed) workers case result of Right conn -> do - mapM_ cancel remaining + cancelAndCloseWorkers remaining pure (Right conn) Left (ex :: SomeException) -> collectResults remaining (show ex : errs) + +-- | Stop losing dial workers and close any connection that crossed the +-- upgrade finish line before cancellation reached it. +cancelAndCloseWorkers :: [Async Connection] -> IO () +cancelAndCloseWorkers workers = do + mapM_ cancel workers + forM_ workers $ \worker -> do + outcome <- waitCatch worker + case outcome of + Right conn -> muxClose (connSession conn) + `catch` \(_ :: SomeException) -> pure () + Left _ -> pure () diff --git a/src/LibP2P/Transport.hs b/src/LibP2P/Transport.hs index 082a562..7686b1d 100644 --- a/src/LibP2P/Transport.hs +++ b/src/LibP2P/Transport.hs @@ -30,7 +30,13 @@ data Listener = Listener -- | Transport provides dial/listen capabilities for a specific protocol. data Transport = Transport - { transportDial :: !(Multiaddr -> IO RawConnection) -- ^ Dial a remote peer + { transportDial :: !(Multiaddr -> IO RawConnection) + -- ^ Dial a remote peer from an ephemeral local port + , transportDialFrom :: !(Maybe Multiaddr -> Multiaddr -> IO RawConnection) + -- ^ Dial a remote peer, optionally binding the local socket first. + -- Hole punching needs this: the outgoing SYN must leave from the + -- listen port so the NAT mapping matches the address advertised in + -- DCUtR CONNECT (specs/relay/DCUtR simultaneous connect). , transportListen :: !(Multiaddr -> IO Listener) -- ^ Listen for inbound connections , transportCanDial :: !(Multiaddr -> Bool) -- ^ Check if this transport can handle the address } diff --git a/src/LibP2P/Transport/TCP.hs b/src/LibP2P/Transport/TCP.hs index b63316f..a8f8784 100644 --- a/src/LibP2P/Transport/TCP.hs +++ b/src/LibP2P/Transport/TCP.hs @@ -8,6 +8,7 @@ module LibP2P.Transport.TCP , socketToStreamIO ) where +import Control.Exception (onException) import Data.Bits (shiftL, shiftR, (.&.), (.|.)) import qualified Data.ByteString as BS import Data.IP (IPv6, fromHostAddress6, toHostAddress6) @@ -22,7 +23,8 @@ import qualified Network.Socket.ByteString as NSB -- | Create a new TCP transport. newTCPTransport :: IO Transport newTCPTransport = pure $ Transport - { transportDial = tcpDial + { transportDial = tcpDialFrom Nothing + , transportDialFrom = tcpDialFrom , transportListen = tcpListen , transportCanDial = canDialTCP } @@ -54,23 +56,53 @@ multiaddrToHostPort addr = case stripP2P addr of -- | Dial a TCP address by directly constructing a SockAddr from the Multiaddr. -- A trailing /p2p/ component is stripped before connecting; the -- original (unstripped) multiaddr is kept as the connection's remote address. -tcpDial :: Multiaddr -> IO RawConnection -tcpDial addr = case stripP2P addr of - Multiaddr [IP4 w, TCP port] -> do - let hostAddr = NS.tupleToHostAddress (octet 3 w, octet 2 w, octet 1 w, octet 0 w) - sockAddr = NS.SockAddrInet (fromIntegral port) hostAddr - sock <- NS.socket NS.AF_INET NS.Stream NS.defaultProtocol - NS.connect sock sockAddr - mkRawConnection sock addr - Multiaddr [IP6 bs, TCP port] -> do - let ipv6 = bytesToIPv6 bs - hostAddr6 = toHostAddress6 ipv6 - sockAddr = NS.SockAddrInet6 (fromIntegral port) 0 hostAddr6 0 - sock <- NS.socket NS.AF_INET6 NS.Stream NS.defaultProtocol - NS.connect sock sockAddr - mkRawConnection sock addr +-- +-- When a local bind address is supplied the outgoing socket is bound to +-- that address (with SO_REUSEADDR/SO_REUSEPORT) so a hole-punch SYN +-- leaves from the listen port. Ordinary dials pass 'Nothing' and get an +-- ephemeral source port. +tcpDialFrom :: Maybe Multiaddr -> Multiaddr -> IO RawConnection +tcpDialFrom mLocal addr = case stripP2P addr of + Multiaddr [IP4 w, TCP port] -> + connectFrom mLocal addr NS.AF_INET (ipv4SockAddr w port) + Multiaddr [IP6 bs, TCP port] -> + connectFrom mLocal addr NS.AF_INET6 (ipv6SockAddr bs port) _ -> fail "tcpDial: unsupported multiaddr" +-- | Create, optionally bind, and connect a TCP socket. +connectFrom :: Maybe Multiaddr -> Multiaddr -> NS.Family -> NS.SockAddr -> IO RawConnection +connectFrom mLocal remoteAddr family sockAddr = do + sock <- NS.socket family NS.Stream NS.defaultProtocol + (do + enableAddrReuse sock + mapM_ (bindLocal sock) mLocal + NS.connect sock sockAddr + mkRawConnection sock remoteAddr + ) `onException` NS.close sock + +-- | Bind a dial socket to a TCP listen address so the SYN uses that port. +bindLocal :: NS.Socket -> Multiaddr -> IO () +bindLocal sock local = case stripP2P local of + Multiaddr [IP4 w, TCP port] -> NS.bind sock (ipv4SockAddr w port) + Multiaddr [IP6 bs, TCP port] -> NS.bind sock (ipv6SockAddr bs port) + _ -> fail "tcpDialFrom: local bind address is not TCP" + +-- | SO_REUSEADDR + SO_REUSEPORT so a listen socket and a hole-punch +-- dial socket can share the same local port. +enableAddrReuse :: NS.Socket -> IO () +enableAddrReuse sock = do + NS.setSocketOption sock NS.ReuseAddr 1 + NS.setSocketOption sock NS.ReusePort 1 + +ipv4SockAddr :: Word32 -> Word16 -> NS.SockAddr +ipv4SockAddr w port = + NS.SockAddrInet (fromIntegral port) + (NS.tupleToHostAddress (octet 3 w, octet 2 w, octet 1 w, octet 0 w)) + +ipv6SockAddr :: BS.ByteString -> Word16 -> NS.SockAddr +ipv6SockAddr bs port = + NS.SockAddrInet6 (fromIntegral port) 0 (toHostAddress6 (bytesToIPv6 bs)) 0 + -- | Create a RawConnection from a connected socket. -- -- TCP_NODELAY is set on every connection (dialed and accepted), matching @@ -97,7 +129,7 @@ tcpListen (Multiaddr [IP4 w, TCP port]) = do let hostAddr = NS.tupleToHostAddress (octet 3 w, octet 2 w, octet 1 w, octet 0 w) sockAddr = NS.SockAddrInet (fromIntegral port) hostAddr sock <- NS.socket NS.AF_INET NS.Stream NS.defaultProtocol - NS.setSocketOption sock NS.ReuseAddr 1 + enableAddrReuse sock NS.bind sock sockAddr NS.listen sock 256 boundSockAddr <- NS.getSocketName sock @@ -115,7 +147,7 @@ tcpListen (Multiaddr [IP6 bs, TCP port]) = do hostAddr6 = toHostAddress6 ipv6 sockAddr = NS.SockAddrInet6 (fromIntegral port) 0 hostAddr6 0 sock <- NS.socket NS.AF_INET6 NS.Stream NS.defaultProtocol - NS.setSocketOption sock NS.ReuseAddr 1 + enableAddrReuse sock NS.bind sock sockAddr NS.listen sock 256 boundSockAddr <- NS.getSocketName sock diff --git a/test/LibP2P/NAT/DCUtR/UpgradeSpec.hs b/test/LibP2P/NAT/DCUtR/UpgradeSpec.hs index 9838f8b..d701149 100644 --- a/test/LibP2P/NAT/DCUtR/UpgradeSpec.hs +++ b/test/LibP2P/NAT/DCUtR/UpgradeSpec.hs @@ -25,6 +25,7 @@ import LibP2P.NAT , NATConfig (..) , defaultDCUtRUpgradeConfig , defaultNATConfig + , dcutrOwnAddrs , holePunchTargets , registerNATHandlers , upgradeRelayedConnection @@ -168,6 +169,33 @@ spec = do after' <- atomically $ lookupConn (swConnPool swA) pidB fmap (isRelayedAddr . connRemoteAddr) after' `shouldBe` Just False + describe "DCUtR CONNECT addresses" $ do + it "should include the relevant relay's Identify observed address" $ + withCircuitTrio fastConfig $ \c -> do + (unrelatedId, _unrelatedKey) <- mkTestIdentity + let observed = Multiaddr [IP4 0xCB007101, TCP 4001] + stale = Multiaddr [IP4 0xCB007102, TCP 4002] + seedObservedAddr (cTargetSw c) (cRelayId c) observed + seedObservedAddr (cTargetSw c) unrelatedId stale + addrs <- dcutrOwnAddrs (cTargetSw c) (Just (cRelayId c)) + addrs `shouldContain` [observed] + addrs `shouldNotContain` [stale] + + it "should retain private listen addresses alongside an observation" $ + withCircuitTrio fastConfig $ \c -> do + let observed = Multiaddr [IP4 0xCB007101, TCP 4001] + seedObservedAddr (cTargetSw c) (cRelayId c) observed + listen <- switchListenAddrsOf (cTargetSw c) + addrs <- dcutrOwnAddrs (cTargetSw c) (Just (cRelayId c)) + addrs `shouldContain` listen + + it "should fall back to listen addresses when no observed address is known" $ do + (sw, _pid) <- newNode fastConfig + bound <- switchListen sw defaultConnectionGater [loopbackAddr] + addrs <- dcutrOwnAddrs sw Nothing + addrs `shouldBe` bound + switchClose sw + describe "hole punch target selection" $ do it "keeps only public, non-relayed advertised addresses" $ withCircuitTrio fastConfig $ \c -> do @@ -312,6 +340,21 @@ seedListenAddrs sw pid addrs = atomically $ , idSignedPeerRecord = Nothing } +-- | Record how a peer has observed us (Identify observedAddr). +seedObservedAddr :: Switch -> PeerId -> Multiaddr -> IO () +seedObservedAddr sw pid addr = atomically $ + modifyTVar' (swPeerStore sw) (Map.insert pid info) + where + info = IdentifyInfo + { idProtocolVersion = Nothing + , idAgentVersion = Nothing + , idPublicKey = Nothing + , idListenAddrs = [] + , idObservedAddr = Just (toBytes addr) + , idProtocols = [] + , idSignedPeerRecord = Nothing + } + -- | The target's relayed connection back to the dialer. targetRelayConn :: Circuit -> IO Connection targetRelayConn c = do diff --git a/test/LibP2P/Switch/ConnectionLifecycleSpec.hs b/test/LibP2P/Switch/ConnectionLifecycleSpec.hs index 1ccd5fb..1adf70b 100644 --- a/test/LibP2P/Switch/ConnectionLifecycleSpec.hs +++ b/test/LibP2P/Switch/ConnectionLifecycleSpec.hs @@ -97,7 +97,15 @@ mkDummyConnection pid openAction = do -- | Mock transport whose dialed connection records rcClose calls in an IORef. mkClosableMockTransport :: KeyPair -> IORef Bool -> IO Transport mkClosableMockTransport responderKP closedRef = pure Transport - { transportDial = \addr -> do + { transportDial = dialFn + , transportDialFrom = \_ -> dialFn + , transportListen = \_ -> error "mock: listen not supported" + , transportCanDial = \(Multiaddr ps) -> case ps of + (IP4 _ : TCP _ : _) -> True + _ -> False + } + where + dialFn addr = do (streamA, streamB) <- mkMemoryStreamPair let rawConnB = RawConnection { rcStreamIO = streamB @@ -114,11 +122,6 @@ mkClosableMockTransport responderKP closedRef = pure Transport , rcRemoteAddr = addr , rcClose = writeIORef closedRef True } - , transportListen = \_ -> error "mock: listen not supported" - , transportCanDial = \(Multiaddr ps) -> case ps of - (IP4 _ : TCP _ : _) -> True - _ -> False - } -- | Build a TCP node with Switch. mkTCPNode :: IO (Switch, PeerId) @@ -281,6 +284,7 @@ spec = do -- A transport whose canDial check blows up mid-dial addTransport sw Transport { transportDial = \_ -> fail "unreachable" + , transportDialFrom = \_ _ -> fail "unreachable" , transportListen = \_ -> error "mock: listen not supported" , transportCanDial = \_ -> error "boom: canDial exploded" } diff --git a/test/LibP2P/Switch/DialSpec.hs b/test/LibP2P/Switch/DialSpec.hs index 81dfbca..ffd1d32 100644 --- a/test/LibP2P/Switch/DialSpec.hs +++ b/test/LibP2P/Switch/DialSpec.hs @@ -21,6 +21,11 @@ import LibP2P.Switch.Dial , recordBackoff ) import LibP2P.Switch (addTransport, newSwitch, switchClose) +import LibP2P.Switch.ResourceManager + ( ResourceManager (..) + , ResourceScope (..) + , emptyUsage + ) import LibP2P.Switch.Types ( BackoffEntry (..) , ConnState (..) @@ -69,7 +74,15 @@ mkDummyConnection pid = do -- the in-memory stream pair. mkMockDialTransport :: KeyPair -> IO Transport mkMockDialTransport responderKP = pure Transport - { transportDial = \addr -> do + { transportDial = dialFn + , transportDialFrom = \_ -> dialFn + , transportListen = \_ -> error "mock: listen not supported" + , transportCanDial = \(Multiaddr ps) -> case ps of + (IP4 _ : TCP _ : _) -> True + _ -> False + } + where + dialFn addr = do (streamA, streamB) <- mkMemoryStreamPair let rawConnB = RawConnection { rcStreamIO = streamB @@ -86,17 +99,20 @@ mkMockDialTransport responderKP = pure Transport , rcRemoteAddr = addr , rcClose = pure () } - , transportListen = \_ -> error "mock: listen not supported" - , transportCanDial = \(Multiaddr ps) -> case ps of - (IP4 _ : TCP _ : _) -> True - _ -> False - } -- | Create a counting mock transport to verify dial deduplication. -- Records the number of transportDial calls in the IORef. mkCountingMockTransport :: KeyPair -> IORef Int -> IO Transport mkCountingMockTransport responderKP counterRef = pure Transport - { transportDial = \addr -> do + { transportDial = dialFn + , transportDialFrom = \_ -> dialFn + , transportListen = \_ -> error "mock: listen not supported" + , transportCanDial = \(Multiaddr ps) -> case ps of + (IP4 _ : TCP _ : _) -> True + _ -> False + } + where + dialFn addr = do atomicModifyIORef' counterRef (\n -> (n + 1, ())) (streamA, streamB) <- mkMemoryStreamPair let rawConnB = RawConnection @@ -114,20 +130,38 @@ mkCountingMockTransport responderKP counterRef = pure Transport , rcRemoteAddr = addr , rcClose = pure () } - , transportListen = \_ -> error "mock: listen not supported" - , transportCanDial = \(Multiaddr ps) -> case ps of - (IP4 _ : TCP _ : _) -> True - _ -> False - } -- | Create a mock transport that always fails to dial. mkFailingTransport :: IO Transport mkFailingTransport = pure Transport { transportDial = \_ -> fail "connection refused" + , transportDialFrom = \_ _ -> fail "connection refused" , transportListen = \_ -> error "mock: listen not supported" , transportCanDial = \_ -> True } +-- | Create a raw connection whose upgrade fails, recording transport cleanup. +mkUnupgradableTransport :: IORef Int -> IO Transport +mkUnupgradableTransport closeCount = pure Transport + { transportDial = dialFn + , transportDialFrom = \_ -> dialFn + , transportListen = \_ -> error "mock: listen not supported" + , transportCanDial = \_ -> True + } + where + close = atomicModifyIORef' closeCount (\n -> (n + 1, ())) + dialFn addr = pure RawConnection + { rcStreamIO = StreamIO + { streamWrite = const (pure ()) + , streamReadByte = fail "upgrade failed" + , streamReadChunk = const (fail "upgrade failed") + , streamClose = close + } + , rcLocalAddr = Multiaddr [IP4 0x7f000001, TCP 0] + , rcRemoteAddr = addr + , rcClose = close + } + spec :: Spec spec = do describe "Backoff" $ do @@ -291,6 +325,23 @@ spec = do backoffResult <- checkBackoff (swDialBackoffs sw) remotePid backoffResult `shouldBe` Left DialBackoff + it "should close the raw connection and release resources when upgrade fails" $ do + (localPid, localKP) <- mkTestIdentity + (remotePid, _remoteKP) <- mkTestIdentity + sw <- newSwitch localPid localKP + closeCount <- newIORef (0 :: Int) + transport <- mkUnupgradableTransport closeCount + addTransport sw transport + + result <- dial sw remotePid [testAddr] + + case result of + Left _ -> pure () + Right _ -> expectationFailure "expected upgrade to fail" + readIORef closeCount `shouldReturn` 1 + usage <- atomically $ readTVar (rsUsage (rmSystemScope (swResourceMgr sw))) + usage `shouldBe` emptyUsage + it "returns DialPeerIdMismatch when remote identity differs from target" $ do (localPid, localKP) <- mkTestIdentity (remotePid, _remoteKP) <- mkTestIdentity diff --git a/test/LibP2P/Transport/TCPSpec.hs b/test/LibP2P/Transport/TCPSpec.hs index 1562566..c15ee42 100644 --- a/test/LibP2P/Transport/TCPSpec.hs +++ b/test/LibP2P/Transport/TCPSpec.hs @@ -3,7 +3,7 @@ module LibP2P.Transport.TCPSpec (spec) where import Control.Concurrent.Async (concurrently) import Control.Exception (SomeException, try) import qualified Data.ByteString as BS -import Data.Word (Word8) +import Data.Word (Word8, Word16) import LibP2P.Multiaddr (Multiaddr (..), encapsulate, fromText) import LibP2P.Multiaddr.Protocol (Protocol (..)) import LibP2P.MultistreamSelect.Negotiation (StreamIO (..)) @@ -103,6 +103,37 @@ spec = do rcClose serverConn listenerClose listener + describe "Hole-punch dial" $ do + it "should bind the local socket to the listen port when dialing from it" $ do + transport <- newTCPTransport + let Right loopback = fromText "/ip4/127.0.0.1/tcp/0" + listener <- transportListen transport loopback + dest <- transportListen transport loopback + (serverConn, clientConn) <- + concurrently + (listenerAccept dest) + (transportDialFrom transport (Just (listenerAddr listener)) (listenerAddr dest)) + tcpPortOf (rcLocalAddr clientConn) `shouldBe` tcpPortOf (listenerAddr listener) + rcClose clientConn + rcClose serverConn + listenerClose dest + listenerClose listener + + it "should use an ephemeral local port for an ordinary dial" $ do + transport <- newTCPTransport + let Right loopback = fromText "/ip4/127.0.0.1/tcp/0" + listener <- transportListen transport loopback + dest <- transportListen transport loopback + (serverConn, clientConn) <- + concurrently + (listenerAccept dest) + (transportDial transport (listenerAddr dest)) + tcpPortOf (rcLocalAddr clientConn) `shouldNotBe` tcpPortOf (listenerAddr listener) + rcClose clientConn + rcClose serverConn + listenerClose dest + listenerClose listener + describe "Dial failure" $ do it "dial to refused port returns error" $ do -- Use a high port on loopback that's very unlikely to be listening. @@ -120,3 +151,9 @@ spec = do testPeerIdMH :: BS.ByteString testPeerIdMH = BS.pack $ [0x00, 0x24, 0x08, 0x01, 0x12, 0x20] <> replicate 32 0xAB +-- | The TCP port of a multiaddr, if it has one. +tcpPortOf :: Multiaddr -> Maybe Word16 +tcpPortOf (Multiaddr ps) = case [p | TCP p <- ps] of + (p : _) -> Just p + [] -> Nothing +