diff --git a/quinn/src/connection.rs b/quinn/src/connection.rs index 69e90b983..1bb6b8dc8 100644 --- a/quinn/src/connection.rs +++ b/quinn/src/connection.rs @@ -220,6 +220,7 @@ impl FuturesStream for IncomingStreams { } else if let Some(ref e) = conn.error { Err(e.clone()) } else if let Some(x) = conn.inner.accept() { + mem::drop(conn); // Release the lock so clone can take it let stream = BiStream::new(self.0.clone(), x); let stream = if x.directionality() == Directionality::Uni { NewStream::Uni(RecvStream(stream)) @@ -242,7 +243,6 @@ pub enum NewStream { Bi(BiStream), } -#[derive(Clone)] pub struct ConnectionRef(Arc>); impl ConnectionRef { @@ -271,17 +271,26 @@ impl ConnectionRef { incoming_streams_reader: None, finishing: FnvHashMap::default(), error: None, + ref_count: 0, }))) } } +impl Clone for ConnectionRef { + fn clone(&self) -> Self { + self.0.lock().unwrap().ref_count += 1; + Self(self.0.clone()) + } +} + impl Drop for ConnectionRef { fn drop(&mut self) { - if Arc::strong_count(&self.0) == 2 { - // If the driver is alive, it's just it and us, so we'd better shut it down. If it's - // not, we can't do any harm. - let conn = &mut *self.0.lock().unwrap(); - if !conn.inner.is_closed() { + let conn = &mut *self.0.lock().unwrap(); + if let Some(x) = conn.ref_count.checked_sub(1) { + conn.ref_count = x; + if x == 0 && !conn.inner.is_closed() { + // If the driver is alive, it's just it and us, so we'd better shut it down. If it's + // not, we can't do any harm. conn.implicit_close(); } } @@ -314,6 +323,8 @@ pub struct ConnectionInner { finishing: FnvHashMap>>, /// Always set to Some before the connection becomes drained error: Option, + /// Number of live handles that can be used to initiate or handle I/O; excludes the driver + ref_count: usize, } impl ConnectionInner { diff --git a/quinn/src/endpoint.rs b/quinn/src/endpoint.rs index 1eef19728..5f291f604 100644 --- a/quinn/src/endpoint.rs +++ b/quinn/src/endpoint.rs @@ -140,7 +140,7 @@ impl Future for EndpointDriver { } } Ok( - if endpoint.unreferenced && endpoint.connections.is_empty() { + if endpoint.ref_count == 0 && endpoint.connections.is_empty() { Async::Ready(()) } else { Async::NotReady @@ -173,11 +173,8 @@ pub(crate) struct EndpointInner { // Stored to give out clones to new ConnectionInners sender: mpsc::UnboundedSender<(ConnectionHandle, EndpointEvent)>, events: mpsc::UnboundedReceiver<(ConnectionHandle, EndpointEvent)>, - /// Whether only one reference to this endpoint remains - /// - /// We presume the final reference to always be the driver, because otherwise nothing we do will - /// have any effect regardless. - unreferenced: bool, + /// Number of live handles that can be used to initiate or handle I/O; excludes the driver + ref_count: usize, } impl EndpointInner { @@ -364,26 +361,29 @@ impl EndpointRef { incoming_reader: None, driver: None, connections: FnvHashMap::default(), - unreferenced: false, + ref_count: 0, }))) } } impl Clone for EndpointRef { fn clone(&self) -> Self { + self.0.lock().unwrap().ref_count += 1; Self(self.0.clone()) } } impl Drop for EndpointRef { fn drop(&mut self) { - if Arc::strong_count(&self.0) == 2 { - // If the driver is about to be on its own, arrange for it to shut down once the last - // connection is gone. - let endpoint = &mut *self.0.lock().unwrap(); - endpoint.unreferenced = true; - if let Some(task) = endpoint.driver.take() { - task.notify(); + let endpoint = &mut *self.0.lock().unwrap(); + if let Some(x) = endpoint.ref_count.checked_sub(1) { + endpoint.ref_count = x; + if x == 0 { + // If the driver is about to be on its own, ensure it can shut down if the last + // connection is gone. + if let Some(task) = endpoint.driver.take() { + task.notify(); + } } } }