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
2 changes: 1 addition & 1 deletion crates/rbitcoin-net/src/ibd/peer_io.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
52 changes: 45 additions & 7 deletions crates/rbitcoin-net/src/peer.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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.
Expand All @@ -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,
Expand All @@ -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(
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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)) => {
Expand Down
94 changes: 94 additions & 0 deletions crates/rbitcoin-net/src/peer_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}
143 changes: 132 additions & 11 deletions crates/rbitcoin-net/src/peers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<Option<tokio::task::AbortHandle>>,
/// Whole session-task abort — drops reader+writer if the loop is stuck.
session_abort: Mutex<Option<tokio::task::AbortHandle>>,
/// 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<Option<std::net::TcpStream>>,
}

impl LivePeer {
Expand Down Expand Up @@ -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<tokio::task::AbortHandle> {
self.writer_abort
.lock()
.unwrap_or_else(|e| e.into_inner())
.take()
}

fn take_session_abort(&self) -> Option<tokio::task::AbortHandle> {
self.session_abort
.lock()
.unwrap_or_else(|e| e.into_inner())
.take()
}

fn take_tcp_shutdown(&self) -> Option<std::net::TcpStream> {
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()
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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<u64> = {
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;
}
}
Expand Down Expand Up @@ -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]
Expand Down
Loading
Loading