From e2cbb86b1e1b2b22c9de3ce28779aa9942016ea4 Mon Sep 17 00:00:00 2001 From: Nathaniel Ledford Date: Sat, 15 Aug 2026 19:46:56 -0400 Subject: [PATCH] fix(home): guarantee scoped stream cancellation --- src/ui/feed_stream_supervisor.rs | 140 +++++++++++++++++++++++++++++-- 1 file changed, 133 insertions(+), 7 deletions(-) diff --git a/src/ui/feed_stream_supervisor.rs b/src/ui/feed_stream_supervisor.rs index 7b0242c..c5bd1d9 100644 --- a/src/ui/feed_stream_supervisor.rs +++ b/src/ui/feed_stream_supervisor.rs @@ -1,7 +1,7 @@ //! Application-owned lifetime supervisor for the aggregate Home Jetstream source. -use std::sync::Arc; use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::{Arc, Mutex}; use futures_util::StreamExt; use iced::futures::SinkExt; @@ -25,6 +25,7 @@ pub(crate) struct Handle { pub(crate) struct Scope { id: u64, handle: Handle, + cancellation: Arc>, } impl Handle { @@ -32,6 +33,7 @@ impl Handle { Scope { id: self.next_scope_id.fetch_add(1, Ordering::Relaxed) + 1, handle: self.clone(), + cancellation: Arc::new(Mutex::new(tokio_util::sync::CancellationToken::new())), } } } @@ -47,6 +49,7 @@ impl Scope { targets: Arc, initial_cursor: u64, ) -> Result<(), SubmitError> { + let cancellation = tokio_util::sync::CancellationToken::new(); self.handle .commands .try_send(Command::Replace { @@ -55,20 +58,46 @@ impl Scope { target_generation, targets, initial_cursor, + cancellation: cancellation.clone(), }) - .map_err(SubmitError::from) + .map_err(SubmitError::from)?; + let previous = Self::replace_cancellation(&self.cancellation, cancellation); + previous.cancel(); + Ok(()) } pub(crate) fn cancel(&self) -> Result<(), SubmitError> { + Self::current_cancellation(&self.cancellation).cancel(); self.handle .commands .try_send(Command::Cancel { scope_id: self.id }) .map_err(SubmitError::from) } + fn current_cancellation( + cancellation: &Mutex, + ) -> tokio_util::sync::CancellationToken { + cancellation + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + } + + fn replace_cancellation( + cancellation: &Mutex, + next: tokio_util::sync::CancellationToken, + ) -> tokio_util::sync::CancellationToken { + std::mem::replace( + &mut *cancellation + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner), + next, + ) + } } impl Drop for Scope { fn drop(&mut self) { + Self::current_cancellation(&self.cancellation).cancel(); let _ = self .handle .commands @@ -112,6 +141,7 @@ enum Command { target_generation: u64, targets: Arc, initial_cursor: u64, + cancellation: tokio_util::sync::CancellationToken, }, Cancel { scope_id: u64, @@ -174,7 +204,7 @@ fn stream_with_source( loop { tokio::select! { Some(command) = command_rx.recv() => match command { - Command::Replace { scope_id, lifecycle_generation, target_generation, targets, initial_cursor } => { + Command::Replace { scope_id, lifecycle_generation, target_generation, targets, initial_cursor, cancellation } => { if let Some(active) = active.take() { active.cancel().await; } @@ -183,10 +213,25 @@ fn stream_with_source( scope_id, task: Some(tokio::spawn(async move { let mut stream = source.stream(targets, initial_cursor); - while let Some(item) = stream.next().await { - if items_tx.send(Event::Item { scope_id, lifecycle_generation, target_generation, item }).await.is_err() { break; } + while let Some(item) = tokio::select! { + _ = cancellation.cancelled() => None, + item = stream.next() => item, + } { + let event = Event::Item { + scope_id, + lifecycle_generation, + target_generation, + item, + }; + let delivered = tokio::select! { + _ = cancellation.cancelled() => false, + result = items_tx.send(event) => result.is_ok(), + }; + if !delivered { + break; + } } - let _ = items_tx.send(Event::Retired { scope_id }).await; + let _ = items_tx.try_send(Event::Retired { scope_id }); })), }); } @@ -216,8 +261,9 @@ mod tests { use std::time::Duration; use futures_util::{SinkExt, StreamExt}; + use tokio::sync::oneshot; - use super::{Event, stream_with_source}; + use super::{Event, SubmitError, stream_with_source}; use crate::authentication::Did; use crate::feed::{HomeFeedStreamItem, HomeFeedStreamTargets, HomeFeedStreamTransportStatus}; use crate::infrastructure::jetstream::JetstreamHomeFeedEventSource; @@ -301,4 +347,84 @@ mod tests { Some(Event::Retired { scope_id }) if scope_id == scope.id() )); } + + #[allow( + clippy::result_large_err, + reason = "tokio-tungstenite fixes the handshake callback's error type." + )] + #[tokio::test] + async fn dropping_a_scope_cancels_its_runner_when_the_command_queue_is_full() { + let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .expect("loopback listener"); + let port = listener.local_addr().expect("listener address").port(); + let (frames_sent, frames_sent_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let (tcp, _) = listener.accept().await.expect("connection"); + let mut socket = tokio_tungstenite::accept_hdr_async( + tcp, + |_request: &tokio_tungstenite::tungstenite::handshake::server::Request, + mut response: tokio_tungstenite::tungstenite::handshake::server::Response| { + response.headers_mut().insert( + "Sec-WebSocket-Protocol", + "xrpc.v1.json".parse().expect("static header"), + ); + Ok(response) + }, + ) + .await + .expect("handshake"); + for sequence in 1..=64 { + if socket + .send(tokio_tungstenite::tungstenite::Message::Text( + format!( + r#"{{"$type":"message","payload":{{"$type":"network.bsky.jetstream.subscribeEvents#commit","seq":{sequence},"did":"did:plc:account","operation":"create","collection":"app.bsky.graph.follow"}}}}"# + ) + .into(), + )) + .await + .is_err() + { + return; + } + } + let _ = frames_sent.send(()); + let _ = tokio::time::timeout(Duration::from_secs(1), socket.next()).await; + }); + let source = Arc::new( + JetstreamHomeFeedEventSource::for_loopback( + url::Url::parse(&format!("ws://127.0.0.1:{port}")).expect("loopback URL"), + ) + .expect("loopback source"), + ); + let mut events = Box::pin(stream_with_source(Ok(source))); + let Event::Ready(handle) = events.next().await.expect("ready event") else { + panic!("supervisor must start with a handle"); + }; + let active = handle.new_scope(); + let targets = Arc::new(HomeFeedStreamTargets::new( + Did::new("did:plc:account").expect("static account DID"), + Vec::new(), + )); + active + .replace(3, 5, Arc::clone(&targets), 7) + .expect("first plan"); + let _ = events.next().await.expect("connecting event"); + frames_sent_rx.await.expect("server frames"); + tokio::task::yield_now().await; + + let queued = handle.new_scope(); + let full = tokio::time::timeout(Duration::from_secs(1), async { + loop { + if queued.replace(4, 6, Arc::clone(&targets), 7) == Err(SubmitError::Full) { + return; + } + tokio::task::yield_now().await; + } + }) + .await; + assert!(full.is_ok(), "command queue must become saturated"); + drop(active); + server.await.expect("server task"); + } } -- 2.51.2