diff --git a/crates/rbitcoin-net/src/ibd/peer_io.rs b/crates/rbitcoin-net/src/ibd/peer_io.rs index 25c419d8..1193092e 100644 --- a/crates/rbitcoin-net/src/ibd/peer_io.rs +++ b/crates/rbitcoin-net/src/ibd/peer_io.rs @@ -177,7 +177,7 @@ pub(crate) async fn spawn_peer( let stream = TcpStream::connect(addr).await?; let ua = rbitcoin_primitives::rbitcoin_subversion(env!("CARGO_PKG_VERSION"), &[] as &[&str]) .unwrap_or_else(|_| format!("/rbitcoin:{}/", env!("CARGO_PKG_VERSION"))); - let (ver, reader, writer, _wire) = connect_and_handshake_timed( + let (ver, reader, writer, _wire, _tcp_shutdown) = connect_and_handshake_timed( HANDSHAKE_TIMEOUT, stream, magic, diff --git a/crates/rbitcoin-net/src/peer.rs b/crates/rbitcoin-net/src/peer.rs index d19db172..c1581673 100644 --- a/crates/rbitcoin-net/src/peer.rs +++ b/crates/rbitcoin-net/src/peer.rs @@ -287,8 +287,17 @@ pub async fn connect_and_handshake( inbound: bool, user_agent: &str, policy: HandshakePolicy<'_>, -) -> Result<(VersionMessage, V2Reader, V2Writer, crate::v2::WireBytes), NetError> { - let (mut reader, mut writer, wire) = open_v2(stream, magic, inbound).await?; +) -> Result< + ( + VersionMessage, + V2Reader, + V2Writer, + crate::v2::WireBytes, + std::net::TcpStream, + ), + NetError, +> { + let (mut reader, mut writer, wire, tcp_shutdown) = open_v2(stream, magic, inbound).await?; let their_version = application_handshake( &mut reader, &mut writer, @@ -301,7 +310,7 @@ pub async fn connect_and_handshake( policy, ) .await?; - Ok((their_version, reader, writer, wire)) + Ok((their_version, reader, writer, wire, tcp_shutdown)) } /// Core VERSION/VERACK bound: 60s from TCP connect/accept. Timeout drops the stream. @@ -315,7 +324,16 @@ pub(crate) async fn inbound_connect_and_handshake( start_height: i32, user_agent: &str, policy: HandshakePolicy<'_>, -) -> Result<(VersionMessage, V2Reader, V2Writer, crate::v2::WireBytes), NetError> { +) -> Result< + ( + VersionMessage, + V2Reader, + V2Writer, + crate::v2::WireBytes, + std::net::TcpStream, + ), + NetError, +> { connect_and_handshake_timed( HANDSHAKE_TIMEOUT, stream, @@ -340,7 +358,16 @@ pub(crate) async fn connect_and_handshake_timed( inbound: bool, user_agent: &str, policy: HandshakePolicy<'_>, -) -> Result<(VersionMessage, V2Reader, V2Writer, crate::v2::WireBytes), NetError> { +) -> Result< + ( + VersionMessage, + V2Reader, + V2Writer, + crate::v2::WireBytes, + std::net::TcpStream, + ), + NetError, +> { tokio::time::timeout( limit, connect_and_handshake( @@ -411,7 +438,7 @@ async fn run_feeler_inner( start_height: i32, user_agent: &str, ) -> Result<(), NetError> { - let (mut reader, mut writer, _wire) = open_v2(stream, magic, false).await?; + let (mut reader, mut writer, _wire, _tcp_shutdown) = open_v2(stream, magic, false).await?; let services = local_service_flags(); let now = SystemTime::now() .duration_since(UNIX_EPOCH) @@ -669,6 +696,9 @@ pub async fn peer_session_with( } } }); + if let Some(s) = meta.session.as_ref() { + s.set_writer_abort(writer_task.abort_handle()); + } if let Some(s) = meta.session.as_ref() { let _ = maybe_queue_addrfetch_getaddr(&out_tx, s); @@ -937,7 +967,15 @@ pub async fn peer_session_with( frame = read_v2_frame(&mut reader, magic) => { let frame = match frame { Ok(f) => f, - Err(NetError::Io(e)) if e.kind() == std::io::ErrorKind::UnexpectedEof => { + Err(NetError::Io(e)) + if matches!( + e.kind(), + std::io::ErrorKind::UnexpectedEof + | std::io::ErrorKind::ConnectionReset + | std::io::ErrorKind::BrokenPipe + | std::io::ErrorKind::ConnectionAborted + ) => + { return Ok(()); } Err(NetError::MessageTooLarge(n)) => { diff --git a/crates/rbitcoin-net/src/peer_tests.rs b/crates/rbitcoin-net/src/peer_tests.rs index 3feb215e..b203e338 100644 --- a/crates/rbitcoin-net/src/peer_tests.rs +++ b/crates/rbitcoin-net/src/peer_tests.rs @@ -5573,3 +5573,97 @@ async fn tip_burst_past_broadcast_capacity_still_syncs_peer() { nb.shutdown().await; let _ = std::fs::remove_dir_all(dir); } + +/// Core `disconnect_nodes` waits ≤5s for the far side's `getpeerinfo` to drop us. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn disconnect_clears_far_side_getpeerinfo_within_5s() { + use crate::P2PNode; + use std::time::Duration; + + let n = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos(); + let dir = std::env::temp_dir().join(format!("rbitcoin-disc-far-{n}")); + std::fs::create_dir_all(dir.join("a")).unwrap(); + std::fs::create_dir_all(dir.join("b")).unwrap(); + let qa = Query::open_or_create(dir.join("a/store")).unwrap(); + let qb = Query::open_or_create(dir.join("b/store")).unwrap(); + let params = ChainParams::regtest(); + let mut na = P2PNode::start_with_agent( + "127.0.0.1:0".parse().unwrap(), + qa, + params.clone(), + Milestone::NONE, + "/rbitcoin:0.1.0(testnode0)/".into(), + crate::DEFAULT_MAX_INBOUND, + ) + .await + .unwrap(); + let nb = P2PNode::start_with_agent( + "127.0.0.1:0".parse().unwrap(), + qb, + params, + Milestone::NONE, + "/rbitcoin:0.1.0(testnode1)/".into(), + crate::DEFAULT_MAX_INBOUND, + ) + .await + .unwrap(); + + na.follow_from(nb.local_addr).await.unwrap(); + let mut linked = false; + for _ in 0..100 { + let a_sees = na + .peers + .snapshot() + .iter() + .any(|p| p.subver.contains("testnode1")); + let b_sees = nb + .peers + .snapshot() + .iter() + .any(|p| p.subver.contains("testnode0")); + if a_sees && b_sees { + linked = true; + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert!(linked, "both sides must list each other before disconnect"); + + let peer_id = na + .peers + .snapshot() + .into_iter() + .find(|p| p.subver.contains("testnode1")) + .map(|p| p.id) + .expect("outbound peer id"); + assert!(na.peers.disconnect_id(peer_id)); + assert!( + na.peers.snapshot().is_empty(), + "local getpeerinfo clears immediately" + ); + + let mut far_clear = false; + for _ in 0..100 { + if !nb + .peers + .snapshot() + .iter() + .any(|p| p.subver.contains("testnode0")) + { + far_clear = true; + break; + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + assert!( + far_clear, + "far side getpeerinfo must drop us within 5s (Core disconnect_nodes)" + ); + + na.shutdown().await; + nb.shutdown().await; + let _ = std::fs::remove_dir_all(dir); +} diff --git a/crates/rbitcoin-net/src/peers.rs b/crates/rbitcoin-net/src/peers.rs index 1d694cf0..e7aeecd3 100644 --- a/crates/rbitcoin-net/src/peers.rs +++ b/crates/rbitcoin-net/src/peers.rs @@ -129,6 +129,13 @@ pub struct LivePeer { connected_at: AtomicU64, /// Skip INV for mempool txs with `accept_gen < floor` (post-verack privacy). inv_gen_floor: AtomicU64, + /// Writer-task abort — FIN via dropping the write half. + writer_abort: Mutex>, + /// Whole session-task abort — drops reader+writer if the loop is stuck. + session_abort: Mutex>, + /// Cloned std TCP fd for `Shutdown::Both` on `disconnectnode` so the far + /// side sees EOF even if our session task is mid-frame. + tcp_shutdown: Mutex>, } impl LivePeer { @@ -196,6 +203,43 @@ impl LivePeer { self.stop.store(true, Ordering::SeqCst); } + pub fn set_writer_abort(&self, handle: tokio::task::AbortHandle) { + *self.writer_abort.lock().unwrap_or_else(|e| e.into_inner()) = Some(handle); + } + + pub fn set_session_abort(&self, handle: tokio::task::AbortHandle) { + *self.session_abort.lock().unwrap_or_else(|e| e.into_inner()) = Some(handle); + } + + pub fn attach_tcp_shutdown(&self, stream: std::net::TcpStream) { + *self.tcp_shutdown.lock().unwrap_or_else(|e| e.into_inner()) = Some(stream); + } + + fn take_writer_abort(&self) -> Option { + self.writer_abort + .lock() + .unwrap_or_else(|e| e.into_inner()) + .take() + } + + fn take_session_abort(&self) -> Option { + self.session_abort + .lock() + .unwrap_or_else(|e| e.into_inner()) + .take() + } + + fn take_tcp_shutdown(&self) -> Option { + self.tcp_shutdown + .lock() + .unwrap_or_else(|e| e.into_inner()) + .take() + } + + fn clear_out_tx(&self) { + *self.out_tx.lock().unwrap_or_else(|e| e.into_inner()) = None; + } + pub fn note_failed_cmpct(&self, hash: BlockHash) { self.failed_cmpct .lock() @@ -1015,6 +1059,9 @@ impl PeerHub { out_tx: Mutex::new(None), connected_at: AtomicU64::new(connected_at), inv_gen_floor: AtomicU64::new(0), + writer_abort: Mutex::new(None), + session_abort: Mutex::new(None), + tcp_shutdown: Mutex::new(None), }); // Handshake already exchanged version + verack (+ maybe ping). peer.note_recv("version", 100); @@ -1170,12 +1217,27 @@ impl PeerHub { } pub fn disconnect_id(&self, id: u64) -> bool { - if let Some(p) = self.get(id) { - p.request_disconnect(); - true - } else { - false + let Some(p) = self.get(id) else { + return false; + }; + p.request_disconnect(); + // Hard-close TCP first so the far side's read sees EOF inside the + // Core `disconnect_nodes` 5s wait (`mempool_reorg`). + if let Some(s) = p.take_tcp_shutdown() { + let _ = s.shutdown(std::net::Shutdown::Both); } + // Drop writer channel + abort writer/session so local halves tear down + // even if the read loop is mid-frame. + p.clear_out_tx(); + if let Some(h) = p.take_writer_abort() { + h.abort(); + } + if let Some(h) = p.take_session_abort() { + h.abort(); + } + // Drop from getpeerinfo before teardown finishes. + self.unregister(id); + true } /// Core `AttemptToEvictConnection`: disconnect one unprotected inbound. @@ -1207,11 +1269,16 @@ impl PeerHub { } pub fn disconnect_addr(&self, addr: SocketAddr) -> bool { - let g = self.live.read().unwrap_or_else(|e| e.into_inner()); + let ids: Vec = { + let g = self.live.read().unwrap_or_else(|e| e.into_inner()); + g.values() + .filter(|p| p.addr == addr) + .map(|p| p.id) + .collect() + }; let mut n = 0usize; - for p in g.values() { - if p.addr == addr { - p.request_disconnect(); + for id in ids { + if self.disconnect_id(id) { n += 1; } } @@ -1306,8 +1373,62 @@ mod tests { assert!(snap[0].bytesrecv_per_msg.get("pong").copied().unwrap() >= 29); assert!(hub.disconnect_id(0)); assert!(p.stop.load(Ordering::SeqCst)); - hub.unregister(0); - assert!(hub.snapshot().is_empty()); + // disconnectnode must clear getpeerinfo immediately (mempool_reorg + // disconnect_nodes waits ≤5s on the far side seeing us gone). + assert!( + hub.snapshot().is_empty(), + "disconnect_id must unregister before the session task exits" + ); + } + + #[test] + fn disconnect_id_shuts_down_tcp_so_far_side_sees_eof() { + use std::io::{Read, Write}; + use std::net::TcpListener; + use std::time::{Duration, Instant}; + + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let addr = listener.local_addr().unwrap(); + let mut local = std::net::TcpStream::connect(addr).unwrap(); + let (mut far, _) = listener.accept().unwrap(); + far.set_read_timeout(Some(Duration::from_millis(200))) + .unwrap(); + + let hub = PeerHub::new(); + let a = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 18444); + let b = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 18445); + let p = hub.register( + a, + b, + &ver("/rbitcoin:0.1.0(testnode0)/"), + false, + PeerConnType::OutboundFullRelay, + ); + let killer = local.try_clone().unwrap(); + p.attach_tcp_shutdown(killer); + + assert!(hub.disconnect_id(0)); + + let start = Instant::now(); + let mut buf = [0u8; 1]; + let n = far.read(&mut buf); + assert!( + start.elapsed() < Duration::from_secs(1), + "far side must observe close promptly, took {:?}", + start.elapsed() + ); + match n { + Ok(0) => {} + Err(e) + if matches!( + e.kind(), + std::io::ErrorKind::ConnectionReset + | std::io::ErrorKind::ConnectionAborted + | std::io::ErrorKind::UnexpectedEof + ) => {} + other => panic!("expected EOF/reset after disconnect_id, got {other:?}"), + } + let _ = local.write(&[1]); } #[test] diff --git a/crates/rbitcoin-net/src/service.rs b/crates/rbitcoin-net/src/service.rs index d1a2d1fa..0886cbaf 100644 --- a/crates/rbitcoin-net/src/service.rs +++ b/crates/rbitcoin-net/src/service.rs @@ -150,12 +150,14 @@ impl P2PNode { let peers = dial_peers.clone(); let ua = dial_ua.clone(); let live = dial_live.clone(); + let (ah_tx, ah_rx) = tokio::sync::oneshot::channel::(); let h = tokio::spawn(async move { - let _ = run_outbound_session( - req.addr, magic, local_addr, hub, peers, ua, live, req.typ, + let _ = run_outbound_session_with_abort( + req.addr, magic, local_addr, hub, peers, ua, live, req.typ, ah_rx, ) .await; }); + let _ = ah_tx.send(h.abort_handle()); push_session_task(&sessions_dial, h); } }); @@ -406,33 +408,40 @@ fn spawn_inbound_accept( Err(_) => our, }; let sessions = session_tasks.clone(); + let (ah_tx, ah_rx) = + tokio::sync::oneshot::channel::(); let h = tokio::spawn(async move { let _session_slot = permit; - let (ver, reader, writer, wire) = match inbound_connect_and_handshake( - stream, - magic, - our, - peer_addr, - height, - &ua, - HandshakePolicy { - hub: Some(hub.as_ref()), - peers: Some(peers.as_ref()), - conn_type: PeerConnType::Inbound, - }, - ) - .await - { - Ok(x) => x, - Err(e) => { - rbitcoin_log::debug!( - "p2p: inbound handshake {peer_addr} failed: {e}" - ); - return; - } - }; + let (ver, reader, writer, wire, tcp_shutdown) = + match inbound_connect_and_handshake( + stream, + magic, + our, + peer_addr, + height, + &ua, + HandshakePolicy { + hub: Some(hub.as_ref()), + peers: Some(peers.as_ref()), + conn_type: PeerConnType::Inbound, + }, + ) + .await + { + Ok(x) => x, + Err(e) => { + rbitcoin_log::debug!( + "p2p: inbound handshake {peer_addr} failed: {e}" + ); + return; + } + }; let sess = peers.register(peer_addr, bind, &ver, true, PeerConnType::Inbound); + sess.attach_tcp_shutdown(tcp_shutdown); + if let Ok(ah) = ah_rx.await { + sess.set_session_abort(ah); + } if let Some(mp) = hub.mempool() { sess.set_inv_gen_floor(mp.next_accept_gen()); } @@ -446,6 +455,7 @@ fn spawn_inbound_accept( let _ = peer_session_with(reader, writer, magic, hub, tip_rx, meta).await; peers.unregister(id); }); + let _ = ah_tx.send(h.abort_handle()); push_session_task(&sessions, h); } Ok(Err(_)) => break, @@ -506,7 +516,7 @@ async fn prepare_outbound_session( }, ) .await; - let (ver, reader, writer, wire) = match handshake { + let (ver, reader, writer, wire, tcp_shutdown) = match handshake { Ok(x) => x, Err(e) => { peers.unregister(provisional_id); @@ -515,6 +525,7 @@ async fn prepare_outbound_session( }; peers.unregister(provisional_id); let sess = peers.register_with_id(provisional_id, peer, bind, &ver, false, typ); + sess.attach_tcp_shutdown(tcp_shutdown); if let Some(mp) = hub.mempool() { sess.set_inv_gen_floor(mp.next_accept_gen()); } @@ -554,7 +565,7 @@ async fn run_prepared_outbound(prepared: PreparedOutbound) -> Result<(), NetErro out } -async fn run_outbound_session( +async fn run_outbound_session_with_abort( peer: SocketAddr, magic: Magic, local: SocketAddr, @@ -563,6 +574,7 @@ async fn run_outbound_session( user_agent: String, follow_live: Arc, typ: PeerConnType, + ah_rx: tokio::sync::oneshot::Receiver, ) -> Result<(), NetError> { if typ == PeerConnType::Feeler { let stream = TcpStream::connect(peer).await?; @@ -572,6 +584,9 @@ async fn run_outbound_session( let prepared = prepare_outbound_session(peer, magic, local, hub, peers, user_agent, follow_live, typ) .await?; + if let Ok(ah) = ah_rx.await { + prepared.sess.set_session_abort(ah); + } run_prepared_outbound(prepared).await } diff --git a/crates/rbitcoin-net/src/v2.rs b/crates/rbitcoin-net/src/v2.rs index 0dae93e9..e9e94d6b 100644 --- a/crates/rbitcoin-net/src/v2.rs +++ b/crates/rbitcoin-net/src/v2.rs @@ -448,14 +448,21 @@ fn map_protocol_error(e: ProtocolError) -> NetError { /// Complete BIP324 handshake on a connected TCP stream; return split encrypted halves. /// +/// The fourth value is a cloned std TCP handle for [`std::net::TcpStream::shutdown`] +/// on `disconnectnode` (far-side EOF without waiting on our session task). +/// /// Not cancellation-safe (BIP324 handshake). Callers should not wrap this in /// `select!` without a dedicated task. pub async fn open_v2( stream: TcpStream, magic: Magic, inbound: bool, -) -> Result<(V2Reader, V2Writer, WireBytes), NetError> { +) -> Result<(V2Reader, V2Writer, WireBytes, std::net::TcpStream), NetError> { let _ = stream.set_nodelay(true); + let std = stream.into_std().map_err(NetError::Io)?; + std.set_nonblocking(true).map_err(NetError::Io)?; + let tcp_shutdown = std.try_clone().map_err(NetError::Io)?; + let stream = TcpStream::from_std(std).map_err(NetError::Io)?; let role = if inbound { Role::Responder } else { @@ -477,7 +484,12 @@ pub async fn open_v2( .await .map_err(map_protocol_error)?; let (r, w) = protocol.into_split(); - Ok((V2SessionReader::from_protocol_reader(r), w, wire)) + Ok(( + V2SessionReader::from_protocol_reader(r), + w, + wire, + tcp_shutdown, + )) } /// Encrypt and send one application message.