diff --git a/crates/video-streamer/Cargo.toml b/crates/video-streamer/Cargo.toml index 857b3aecc..2d926a888 100644 --- a/crates/video-streamer/Cargo.toml +++ b/crates/video-streamer/Cargo.toml @@ -53,3 +53,4 @@ workspace = true [[bench]] name = "vpx_reencode" harness = false +required-features = ["bench"] diff --git a/crates/video-streamer/src/decoder.rs b/crates/video-streamer/src/decoder.rs new file mode 100644 index 000000000..b8fce099e --- /dev/null +++ b/crates/video-streamer/src/decoder.rs @@ -0,0 +1,55 @@ +use anyhow::Context as _; +use cadeau::xmf::vpx::{VpxCodec, VpxDecoder, VpxImage}; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) struct Dimensions { + pub width: u32, + pub height: u32, +} + +pub(crate) struct DecodedFrame<'decoder> { + pub image: VpxImage<'decoder>, + pub dimensions: Dimensions, +} + +pub(crate) struct InputDecoder { + codec: VpxCodec, + threads: u32, + decoder: Option, +} + +impl InputDecoder { + pub(crate) fn new(codec: VpxCodec, threads: u32) -> Self { + Self { + codec, + threads, + decoder: None, + } + } + + pub(crate) fn decode<'decoder>(&'decoder mut self, data: &[u8]) -> anyhow::Result> { + if self.decoder.is_none() { + self.decoder = Some( + VpxDecoder::builder() + .threads(self.threads) + .width(0) + .height(0) + .codec(self.codec) + .build()?, + ); + } + + let decoder = self.decoder.as_mut().context("input decoder is missing")?; + decoder.decode(data)?; + let image = decoder.next_frame()?; + let dimensions = Dimensions { + width: image.width(), + height: image.height(), + }; + anyhow::ensure!( + dimensions.width > 0 && dimensions.height > 0, + "decoder returned invalid frame dimensions" + ); + Ok(DecodedFrame { image, dimensions }) + } +} diff --git a/crates/video-streamer/src/lib.rs b/crates/video-streamer/src/lib.rs index e568689a0..9ffe2a900 100644 --- a/crates/video-streamer/src/lib.rs +++ b/crates/video-streamer/src/lib.rs @@ -25,7 +25,11 @@ macro_rules! perf_debug { pub mod config; pub mod debug; +mod decoder; +mod normalizer; +mod protocol; pub mod reopenable; +mod session; pub(crate) mod streamer; #[macro_use] @@ -39,6 +43,8 @@ pub use streamer::reopenable_file::ReOpenableFile; pub use streamer::signal_writer::SignalWriter; #[rustfmt::skip] pub use streamer::webm_stream; +#[rustfmt::skip] +pub use session::{RecordingClip, RecordingEvent, RecordingSource, SessionConfig, StartAt, stream_session}; #[cfg(feature = "bench")] pub mod bench_support; diff --git a/crates/video-streamer/src/normalizer/mod.rs b/crates/video-streamer/src/normalizer/mod.rs new file mode 100644 index 000000000..6a6bbcbfc --- /dev/null +++ b/crates/video-streamer/src/normalizer/mod.rs @@ -0,0 +1,1019 @@ +use std::io::{self, SeekFrom, Write}; +use std::pin::Pin; +use std::task::{Context as TaskContext, Poll}; +use std::time::{Duration, Instant}; + +use anyhow::Context; +use bytes::{Bytes, BytesMut}; +use cadeau::xmf::vpx::{VpxCodec, VpxEncoder, VpxEncoderPreset, VpxImage}; +use ebml_iterable::error::TagIteratorError; +use ebml_iterable::{PositionedTag, TagDecoder}; +use futures_util::{Stream, StreamExt}; +use tokio::sync::mpsc; +use webm_iterable::matroska_spec::{Master, MatroskaSpec, SimpleBlock}; +use webm_iterable::{WebmWriter, WriteOptions}; + +use crate::decoder::{Dimensions, InputDecoder}; +use crate::session::{RecordingClip, RecordingEvent, SessionConfig, StartAt}; +use crate::streamer::block_tag::{VideoBlock, is_vpx_key_frame}; + +const OUTPUT_CHANNEL_CAPACITY: usize = 4; +const OUTPUT_CHUNK_SIZE: usize = 64 * 1024; +const INPUT_CHANNEL_CAPACITY: usize = 1; +const INPUT_CHUNK_SIZE: usize = 64 * 1024; +const MAX_TAG_PAYLOAD_BYTES: usize = 64 * 1024 * 1024; +const MAX_INPUT_BUFFER_BYTES: usize = MAX_TAG_PAYLOAD_BYTES + 16; +const OUTPUT_BITRATE: u32 = 256 * 1024; +const VPX_EFLAG_FORCE_KF: u32 = 0x0000_0001; +const WEBM_TIMESTAMP_SCALE_NS: u64 = 1_000_000; +const MAX_WEBM_BLOCK_TIMESTAMP: u64 = 32_767; +const MAX_CONSECUTIVE_FRAME_SKIPS: u32 = 1; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) struct SegmentInfo { + pub sequence: u64, + pub width: u32, + pub height: u32, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub(crate) enum SegmentEvent { + Begin(SegmentInfo), + Data(Bytes), + End, +} + +pub(crate) struct NormalizedSession { + receiver: mpsc::Receiver>, + supervisor: Option>, +} + +impl Stream for NormalizedSession { + type Item = anyhow::Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll> { + self.receiver.poll_recv(cx) + } +} + +impl NormalizedSession { + pub(crate) async fn shutdown(mut self) -> anyhow::Result<()> { + self.receiver.close(); + let supervisor = self.supervisor.take().context("normalizer supervisor is missing")?; + supervisor.await.context("normalizer supervisor failed") + } +} + +impl Drop for NormalizedSession { + fn drop(&mut self) { + if let Some(supervisor) = self.supervisor.take() { + supervisor.abort(); + } + } +} + +#[cfg(test)] +pub(crate) fn test_session(stream: S) -> NormalizedSession +where + S: Stream> + Send + 'static, +{ + let (sender, receiver) = mpsc::channel(OUTPUT_CHANNEL_CAPACITY); + let supervisor = tokio::spawn(async move { + tokio::pin!(stream); + loop { + tokio::select! { + event = stream.next() => { + let Some(event) = event else { break }; + if sender.send(event).await.is_err() { + break; + } + } + () = sender.closed() => break, + } + } + }); + NormalizedSession { + receiver, + supervisor: Some(supervisor), + } +} + +pub(crate) fn normalize(source: S, config: SessionConfig) -> NormalizedSession +where + S: Stream> + Send + 'static, +{ + let (output_sender, output_receiver) = mpsc::channel(OUTPUT_CHANNEL_CAPACITY); + let (input_sender, input_receiver) = mpsc::channel(INPUT_CHANNEL_CAPACITY); + + let supervisor = tokio::spawn(async move { + let worker_sender = output_sender.clone(); + let mut worker = tokio::task::spawn_blocking(move || normalize_events(input_receiver, worker_sender, config)); + let mut forward = Box::pin(async move { + tokio::pin!(source); + while let Some(event) = source.next().await { + if input_sender.send(event).await.is_err() { + break; + } + } + }); + + tokio::select! { + result = &mut worker => publish_worker_result(result, &output_sender).await, + () = output_sender.closed() => { + drop(forward); + let _ = worker.await; + } + () = &mut forward => { + drop(forward); + publish_worker_result(worker.await, &output_sender).await; + } + }; + }); + + NormalizedSession { + receiver: output_receiver, + supervisor: Some(supervisor), + } +} + +async fn publish_worker_result( + result: Result, tokio::task::JoinError>, + sender: &mpsc::Sender>, +) { + let error = match result { + Ok(Ok(())) => return, + Ok(Err(error)) => error.context("session normalization failed"), + Err(error) => anyhow::Error::new(error).context("normalizer worker failed"), + }; + let _ = sender.send(Err(error)).await; +} + +fn normalize_events( + mut receiver: mpsc::Receiver>, + sender: mpsc::Sender>, + config: SessionConfig, +) -> anyhow::Result<()> { + let mut phase = SessionPhase::AwaitClip; + let mut next_segment_sequence = 0; + + while let Some(event) = receiver.blocking_recv() { + match event.context("recording source failed")? { + RecordingEvent::ClipStarted { + sequence, + start_at, + clip, + } => { + anyhow::ensure!( + matches!(phase, SessionPhase::AwaitClip), + "clip {sequence} started before the previous clip ended" + ); + let mut clip_normalizer = + ClipNormalizer::new(sequence, start_at, clip, sender.clone(), config, next_segment_sequence)?; + clip_normalizer.scan_available()?; + phase = SessionPhase::InClip(Box::new(clip_normalizer)); + } + RecordingEvent::DataAvailable => { + let SessionPhase::InClip(clip) = &mut phase else { + anyhow::bail!("data availability arrived outside a clip"); + }; + clip.scan_available()?; + } + RecordingEvent::CaughtUp => { + let SessionPhase::InClip(clip) = &mut phase else { + anyhow::bail!("caught-up arrived outside a clip"); + }; + clip.caught_up()?; + } + RecordingEvent::ClipEnded => { + let SessionPhase::InClip(current) = std::mem::replace(&mut phase, SessionPhase::AwaitClip) else { + anyhow::bail!("clip end arrived outside a clip"); + }; + next_segment_sequence = (*current).finish()?; + } + RecordingEvent::SessionEnded => { + anyhow::ensure!( + matches!(phase, SessionPhase::AwaitClip), + "session ended before the active clip ended" + ); + phase = SessionPhase::Ended; + break; + } + } + } + + anyhow::ensure!( + matches!(phase, SessionPhase::Ended), + "recording source ended before the session end event" + ); + Ok(()) +} + +enum SessionPhase { + AwaitClip, + InClip(Box), + Ended, +} + +#[derive(Clone, Copy)] +struct SourceVideo { + track: u64, + codec: VpxCodec, +} + +#[derive(Default)] +struct TrackEntryState { + track: Option, + track_type: Option, + codec_id: Option, +} + +struct PendingFrame { + data: Vec, + timestamp: u64, + codec: VpxCodec, + key_frame: bool, +} + +enum ClipPhase { + History(HistoryPolicy), + Live, +} + +enum HistoryPolicy { + EmitAll, + KeepLatestGop, +} + +#[derive(Clone, Copy)] +struct ReplayPoint { + block_offset: u64, + cluster_timestamp: u64, +} + +struct PendingBlockGroup { + offset: u64, + block: Option>, +} + +struct ClipNormalizer { + clip_sequence: u64, + clip: RecordingClip, + reader_head: u64, + decoder: TagDecoder, + input: BytesMut, + source_video: Option, + track_entry: Option, + pending_block_group: Option, + cluster_timestamp: Option, + timestamp_scale_ns: u64, + phase: ClipPhase, + replay_point: Option, + complete_boundary: u64, + input_decoder: Option, + output_segment: Option, + next_segment_sequence: u64, + processing_time: Duration, + first_frame_timestamp: Option, + frames_since_last_encode: u32, + sender: mpsc::Sender>, + config: SessionConfig, +} + +impl ClipNormalizer { + fn new( + clip_sequence: u64, + start_at: StartAt, + mut clip: RecordingClip, + sender: mpsc::Sender>, + config: SessionConfig, + next_segment_sequence: u64, + ) -> anyhow::Result { + let reader_head = clip.seek(SeekFrom::Start(0))?; + anyhow::ensure!(reader_head == 0, "recording clip did not seek to its beginning"); + let phase = match start_at { + StartAt::Beginning => ClipPhase::History(HistoryPolicy::EmitAll), + StartAt::LiveEdge => ClipPhase::History(HistoryPolicy::KeepLatestGop), + }; + Ok(Self { + clip_sequence, + clip, + reader_head, + decoder: new_decoder(), + input: BytesMut::new(), + source_video: None, + track_entry: None, + pending_block_group: None, + cluster_timestamp: None, + timestamp_scale_ns: WEBM_TIMESTAMP_SCALE_NS, + phase, + replay_point: None, + complete_boundary: reader_head, + input_decoder: None, + output_segment: None, + next_segment_sequence, + processing_time: Duration::ZERO, + first_frame_timestamp: None, + frames_since_last_encode: 0, + sender, + config, + }) + } + + fn scan_available(&mut self) -> anyhow::Result<()> { + let mut process_frame = Self::process_frame; + self.scan_available_with(&mut process_frame) + } + + fn scan_available_with(&mut self, process_frame: &mut F) -> anyhow::Result<()> + where + F: FnMut(&mut Self, PendingFrame) -> anyhow::Result<()>, + { + loop { + if self.sender.is_closed() { + return Ok(()); + } + + while let Some(positioned) = self.decoder.decode(&mut self.input)? { + if self.sender.is_closed() { + return Ok(()); + } + self.handle_tag_with(positioned, process_frame)?; + } + + if self.sender.is_closed() { + return Ok(()); + } + if self.input.len() >= MAX_INPUT_BUFFER_BYTES { + anyhow::bail!("recording input exceeds the resource limit"); + } + + let read_limit = (MAX_INPUT_BUFFER_BYTES - self.input.len()).min(INPUT_CHUNK_SIZE); + anyhow::ensure!(read_limit > 0, "recording input cannot make progress"); + let mut buffer = vec![0; read_limit]; + let read = self.clip.read(&mut buffer)?; + if read == 0 { + return Ok(()); + } + self.reader_head = self + .reader_head + .checked_add(u64::try_from(read).context("recording reader position overflow")?) + .context("recording reader position overflow")?; + self.input.extend_from_slice(&buffer[..read]); + } + } + + fn caught_up(&mut self) -> anyhow::Result<()> { + let mut process_frame = Self::process_frame; + self.caught_up_with(&mut process_frame) + } + + fn caught_up_with(&mut self, process_frame: &mut F) -> anyhow::Result<()> + where + F: FnMut(&mut Self, PendingFrame) -> anyhow::Result<()>, + { + let history = match std::mem::replace(&mut self.phase, ClipPhase::Live) { + ClipPhase::History(history) => history, + ClipPhase::Live => anyhow::bail!("clip {} sent caught-up twice", self.clip_sequence), + }; + if matches!(history, HistoryPolicy::KeepLatestGop) + && let Some(replay_point) = self.replay_point + { + self.replay_latest_gop_with(replay_point, process_frame)?; + } + Ok(()) + } + + fn finish(mut self) -> anyhow::Result { + anyhow::ensure!( + matches!(self.phase, ClipPhase::Live), + "clip {} ended before caught-up", + self.clip_sequence + ); + self.scan_available()?; + if self.sender.is_closed() { + return Ok(self.next_segment_sequence); + } + loop { + if self.sender.is_closed() { + return Ok(self.next_segment_sequence); + } + match self.decoder.decode_eof(&mut self.input) { + Ok(Some(positioned)) => self.handle_tag(positioned)?, + Ok(None) if self.decoder.is_finished() => break, + Ok(None) => continue, + Err(TagIteratorError::UnexpectedEOF { .. }) => { + debug!( + clip_sequence = self.clip_sequence, + bytes = self.input.len(), + "Discard incomplete trailing EBML element" + ); + self.input.clear(); + self.pending_block_group = None; + self.track_entry = None; + break; + } + Err(error) => return Err(error.into()), + } + } + + if let Some(segment) = self.output_segment.take() { + segment.finish()?; + } + Ok(self.next_segment_sequence) + } + + fn handle_tag(&mut self, positioned: PositionedTag) -> anyhow::Result<()> { + let mut process_frame = Self::process_frame; + self.handle_tag_with(positioned, &mut process_frame) + } + + fn handle_tag_with( + &mut self, + positioned: PositionedTag, + process_frame: &mut F, + ) -> anyhow::Result<()> + where + F: FnMut(&mut Self, PendingFrame) -> anyhow::Result<()>, + { + let offset = u64::try_from(positioned.offset).context("recording tag offset overflow")?; + match positioned.tag { + MatroskaSpec::TrackEntry(Master::Start) => { + anyhow::ensure!(self.track_entry.is_none(), "nested video track entry"); + self.track_entry = Some(TrackEntryState::default()); + } + MatroskaSpec::TrackEntry(Master::End) => { + let track_entry = self + .track_entry + .take() + .context("track entry end arrived without a start")?; + self.finish_track_entry(track_entry)?; + } + MatroskaSpec::TrackNumber(value) => { + if let Some(track_entry) = &mut self.track_entry { + track_entry.track = Some(value); + } + } + MatroskaSpec::TrackType(value) => { + if let Some(track_entry) = &mut self.track_entry { + track_entry.track_type = Some(value); + } + } + MatroskaSpec::CodecID(value) => { + if let Some(track_entry) = &mut self.track_entry { + track_entry.codec_id = Some(value); + } + } + MatroskaSpec::TimestampScale(value) => { + self.timestamp_scale_ns = value; + } + MatroskaSpec::Cluster(Master::Start) => self.cluster_timestamp = None, + MatroskaSpec::Timestamp(value) => self.cluster_timestamp = Some(value), + MatroskaSpec::BlockGroup(Master::Start) => { + anyhow::ensure!(self.pending_block_group.is_none(), "nested block group"); + self.pending_block_group = Some(PendingBlockGroup { offset, block: None }); + } + MatroskaSpec::Block(data) => { + let group = self + .pending_block_group + .as_mut() + .context("block arrived outside a block group")?; + anyhow::ensure!( + group.block.replace(data).is_none(), + "block group contains multiple blocks" + ); + } + MatroskaSpec::BlockGroup(Master::End) => { + let group = self + .pending_block_group + .take() + .context("block group end arrived without a start")?; + let data = group.block.context("block group does not contain a block")?; + self.handle_block_with( + MatroskaSpec::BlockGroup(Master::Full(vec![MatroskaSpec::Block(data)])), + group.offset, + process_frame, + )?; + self.complete_boundary = decoder_position(&self.decoder)?; + } + MatroskaSpec::SimpleBlock(data) => { + self.handle_block_with(MatroskaSpec::SimpleBlock(data), offset, process_frame)?; + self.complete_boundary = decoder_position(&self.decoder)?; + } + _ => {} + } + Ok(()) + } + + fn finish_track_entry(&mut self, track_entry: TrackEntryState) -> anyhow::Result<()> { + if track_entry.track_type != Some(1) { + return Ok(()); + } + anyhow::ensure!(self.source_video.is_none(), "multiple video tracks are not supported"); + let track = track_entry.track.context("video track number is missing")?; + let codec_id = track_entry.codec_id.context("video codec ID is missing")?; + let codec = match codec_id.as_str() { + "V_VP8" | "vp8" => VpxCodec::VP8, + "V_VP9" | "vp9" => VpxCodec::VP9, + _ => anyhow::bail!("unsupported video codec: {codec_id}"), + }; + self.source_video = Some(SourceVideo { track, codec }); + Ok(()) + } + + fn handle_block_with( + &mut self, + tag: MatroskaSpec, + block_offset: u64, + process_frame: &mut F, + ) -> anyhow::Result<()> + where + F: FnMut(&mut Self, PendingFrame) -> anyhow::Result<()>, + { + let cluster_timestamp = self.cluster_timestamp; + let Some(frame) = self.frame_from_block(tag, cluster_timestamp)? else { + return Ok(()); + }; + + if matches!(&self.phase, ClipPhase::History(HistoryPolicy::KeepLatestGop)) { + if frame.key_frame { + self.replay_point = Some(ReplayPoint { + block_offset, + cluster_timestamp: cluster_timestamp.context("cluster timestamp is missing")?, + }); + } + return Ok(()); + } + + process_frame(self, frame) + } + + fn frame_from_block( + &self, + tag: MatroskaSpec, + cluster_timestamp: Option, + ) -> anyhow::Result> { + let video = self + .source_video + .context("video track header not found before video data")?; + let block = VideoBlock::new(tag, cluster_timestamp, video.codec)?; + if block.track != video.track { + return Ok(None); + } + + let data = block.get_frame()?; + let key_frame = is_vpx_key_frame(&data, video.codec); + let timestamp = scale_timestamp(block.absolute_timestamp()?, self.timestamp_scale_ns)?; + Ok(Some(PendingFrame { + data, + timestamp, + codec: video.codec, + key_frame, + })) + } + + fn process_frame(&mut self, frame: PendingFrame) -> anyhow::Result<()> { + let should_skip_encode = self.should_skip_encode(frame.timestamp); + let processing_started = Instant::now(); + let result = self.process_frame_with_skip(frame, should_skip_encode); + self.processing_time += processing_started.elapsed(); + result + } + + fn process_frame_with_skip(&mut self, frame: PendingFrame, should_skip_encode: bool) -> anyhow::Result<()> { + let input_decoder = self + .input_decoder + .get_or_insert_with(|| InputDecoder::new(frame.codec, self.config.encoder_threads)); + let decoded = input_decoder.decode(&frame.data)?; + let dimensions = decoded.dimensions; + let new_segment = next_segment_info( + self.output_segment.as_ref().map(|segment| segment.dimensions), + dimensions, + self.next_segment_sequence, + ); + + if new_segment.is_none() && should_skip_encode { + self.frames_since_last_encode += 1; + return Ok(()); + } + self.frames_since_last_encode = 0; + + if self.output_segment.is_some() && new_segment.is_some() { + self.output_segment + .take() + .context("missing active output segment")? + .finish()?; + } + + if let Some(info) = new_segment { + self.output_segment = Some(OutputSegment::new(self.sender.clone(), info, self.config)?); + self.next_segment_sequence = self + .next_segment_sequence + .checked_add(1) + .context("segment sequence overflow")?; + } + self.output_segment + .as_mut() + .context("output segment is missing")? + .encode(&decoded.image, frame.timestamp)?; + Ok(()) + } + + fn should_skip_encode(&mut self, timestamp: u64) -> bool { + let first_timestamp = *self.first_frame_timestamp.get_or_insert(timestamp); + let media_advanced_ms = timestamp.saturating_sub(first_timestamp); + let processing_ms = u64::try_from(self.processing_time.as_millis()).unwrap_or(u64::MAX); + + self.config.adaptive_frame_skip + && processing_ms > media_advanced_ms + && self.frames_since_last_encode < MAX_CONSECUTIVE_FRAME_SKIPS + } + + fn replay_latest_gop_with(&mut self, replay_point: ReplayPoint, process_frame: &mut F) -> anyhow::Result<()> + where + F: FnMut(&mut Self, PendingFrame) -> anyhow::Result<()>, + { + if self.sender.is_closed() { + return Ok(()); + } + + let replay_end = self.complete_boundary; + anyhow::ensure!( + replay_point.block_offset <= replay_end, + "replay point is after the complete scan boundary" + ); + if replay_point.block_offset == replay_end { + return Ok(()); + } + + let original_decoder = std::mem::replace(&mut self.decoder, new_decoder()); + let original_input = std::mem::take(&mut self.input); + let original_reader_head = self.reader_head; + let replay_result = (|| { + self.seek_reader(replay_point.block_offset)?; + self.replay_window(replay_point, replay_end, process_frame) + })(); + + self.decoder = original_decoder; + self.input = original_input; + let restore_result = self.seek_reader(original_reader_head); + if let Err(restore_error) = restore_result { + return Err(match replay_result { + Ok(()) => restore_error.context("failed to restore recording reader"), + Err(replay_error) => { + replay_error.context(format!("failed to restore recording reader: {restore_error:#}")) + } + }); + } + replay_result + } + + fn replay_window( + &mut self, + replay_point: ReplayPoint, + replay_end: u64, + process_frame: &mut F, + ) -> anyhow::Result<()> + where + F: FnMut(&mut Self, PendingFrame) -> anyhow::Result<()>, + { + let mut cluster_timestamp = Some(replay_point.cluster_timestamp); + let mut block_group: Option = None; + + loop { + if self.sender.is_closed() { + return Ok(()); + } + + while let Some(positioned) = self.decoder.decode(&mut self.input)? { + if self.sender.is_closed() { + return Ok(()); + } + self.handle_replay_tag( + positioned, + replay_point.block_offset, + &mut cluster_timestamp, + &mut block_group, + process_frame, + )?; + } + + if self.reader_head >= replay_end { + anyhow::ensure!(self.input.is_empty(), "replay endpoint is inside an incomplete element"); + if let Some(group) = block_group.take() { + anyhow::ensure!( + group.offset < replay_end, + "replay endpoint is inside an incomplete block group" + ); + self.process_replay_block_group(group, cluster_timestamp, process_frame)?; + } + return Ok(()); + } + + if self.input.len() >= MAX_INPUT_BUFFER_BYTES { + anyhow::bail!("replay input exceeds the resource limit"); + } + let read_limit = usize::try_from(replay_end - self.reader_head) + .context("replay window is too large")? + .min(INPUT_CHUNK_SIZE) + .min(MAX_INPUT_BUFFER_BYTES - self.input.len()); + anyhow::ensure!(read_limit > 0, "replay input cannot make progress"); + let mut buffer = vec![0; read_limit]; + let read = self.clip.read(&mut buffer)?; + anyhow::ensure!(read > 0, "recording ended before replay boundary"); + self.reader_head = self + .reader_head + .checked_add(u64::try_from(read).context("replay reader position overflow")?) + .context("replay reader position overflow")?; + self.input.extend_from_slice(&buffer[..read]); + } + } + + fn process_replay_block_group( + &mut self, + group: PendingBlockGroup, + cluster_timestamp: Option, + process_frame: &mut F, + ) -> anyhow::Result<()> + where + F: FnMut(&mut Self, PendingFrame) -> anyhow::Result<()>, + { + let data = group.block.context("replay block group does not contain a block")?; + if let Some(frame) = self.frame_from_block( + MatroskaSpec::BlockGroup(Master::Full(vec![MatroskaSpec::Block(data)])), + cluster_timestamp, + )? { + process_frame(self, frame)?; + } + Ok(()) + } + + fn handle_replay_tag( + &mut self, + positioned: PositionedTag, + replay_offset: u64, + cluster_timestamp: &mut Option, + block_group: &mut Option, + process_frame: &mut F, + ) -> anyhow::Result<()> + where + F: FnMut(&mut Self, PendingFrame) -> anyhow::Result<()>, + { + let offset = u64::try_from(positioned.offset) + .context("replay tag offset overflow")? + .checked_add(replay_offset) + .context("replay tag offset overflow")?; + match positioned.tag { + MatroskaSpec::Cluster(Master::Start) => *cluster_timestamp = None, + MatroskaSpec::Timestamp(value) => *cluster_timestamp = Some(value), + MatroskaSpec::BlockGroup(Master::Start) => { + anyhow::ensure!(block_group.is_none(), "nested replay block group"); + *block_group = Some(PendingBlockGroup { offset, block: None }); + } + MatroskaSpec::Block(data) => { + let group = block_group + .as_mut() + .context("replay block arrived outside a block group")?; + anyhow::ensure!( + group.block.replace(data).is_none(), + "replay block group contains multiple blocks" + ); + } + MatroskaSpec::BlockGroup(Master::End) => { + let group = block_group.take().context("replay block group end without a start")?; + self.process_replay_block_group(group, *cluster_timestamp, process_frame)?; + } + MatroskaSpec::SimpleBlock(data) => { + if let Some(frame) = self.frame_from_block(MatroskaSpec::SimpleBlock(data), *cluster_timestamp)? { + process_frame(self, frame)?; + } + } + _ => {} + } + Ok(()) + } + + fn seek_reader(&mut self, position: u64) -> anyhow::Result<()> { + self.clip.seek(SeekFrom::Start(position))?; + self.reader_head = position; + Ok(()) + } +} + +fn new_decoder() -> TagDecoder { + let mut decoder = TagDecoder::new(&[]); + decoder.set_max_allowable_tag_size(Some(MAX_TAG_PAYLOAD_BYTES)); + decoder +} + +fn decoder_position(decoder: &TagDecoder) -> anyhow::Result { + u64::try_from(decoder.position()).context("decoder position overflow") +} + +fn next_segment_info( + current_dimensions: Option, + frame_dimensions: Dimensions, + next_sequence: u64, +) -> Option { + (current_dimensions != Some(frame_dimensions)).then_some(SegmentInfo { + sequence: next_sequence, + width: frame_dimensions.width, + height: frame_dimensions.height, + }) +} + +fn scale_timestamp(value: u64, timestamp_scale_ns: u64) -> anyhow::Result { + let nanoseconds = u128::from(value) + .checked_mul(u128::from(timestamp_scale_ns)) + .context("video timestamp overflow")?; + u64::try_from(nanoseconds / u128::from(WEBM_TIMESTAMP_SCALE_NS)).context("video timestamp is too large") +} + +struct OutputSegment { + info: SegmentInfo, + dimensions: Dimensions, + origin_timestamp: Option, + previous_timestamp: Option, + cluster_timestamp: Option, + encoder: VpxEncoder, + writer: WebmWriter, +} + +impl OutputSegment { + fn new( + sender: mpsc::Sender>, + info: SegmentInfo, + config: SessionConfig, + ) -> anyhow::Result { + send_event(&sender, SegmentEvent::Begin(info))?; + + let encoder = VpxEncoder::builder() + .timebase_num(1) + .timebase_den(1000) + .codec(VpxCodec::VP8) + .width(info.width) + .height(info.height) + .threads(config.encoder_threads) + .bitrate(OUTPUT_BITRATE) + .preset(VpxEncoderPreset::BestPerformance) + .build()?; + let mut writer = WebmWriter::new(EventWriter { sender }); + write_header(&mut writer, info.width, info.height)?; + + Ok(Self { + info, + dimensions: Dimensions { + width: info.width, + height: info.height, + }, + origin_timestamp: None, + previous_timestamp: None, + cluster_timestamp: None, + encoder, + writer, + }) + } + + fn encode(&mut self, image: &VpxImage<'_>, timestamp: u64) -> anyhow::Result<()> { + let origin = *self.origin_timestamp.get_or_insert(timestamp); + let relative_timestamp = timestamp.saturating_sub(origin); + let duration = self + .previous_timestamp + .map_or(30, |previous| timestamp.saturating_sub(previous).max(1)); + self.previous_timestamp = Some(timestamp); + + let cluster_timestamp_expired = self.cluster_timestamp.is_some_and(|cluster_timestamp| { + relative_timestamp.saturating_sub(cluster_timestamp) > MAX_WEBM_BLOCK_TIMESTAMP + }); + let flags = if relative_timestamp == 0 || cluster_timestamp_expired { + VPX_EFLAG_FORCE_KF + } else { + 0 + }; + self.encoder.encode_frame( + image, + i64::try_from(relative_timestamp).context("relative timestamp is too large")?, + usize::try_from(duration).unwrap_or(usize::MAX), + flags, + )?; + self.write_encoded_frames() + } + + fn write_encoded_frames(&mut self) -> anyhow::Result<()> { + let frames = self + .encoder + .packet_iterator() + .filter_map(|packet| packet.frame()) + .map(|frame| { + let timestamp = u64::try_from(frame.pts()).context("encoder returned a negative timestamp")?; + let data = frame.buffer().context("encoder returned a frame without data")?; + Ok((timestamp, data)) + }) + .collect::>>()?; + + for (timestamp, data) in frames { + let is_key_frame = is_vpx_key_frame(&data, VpxCodec::VP8); + anyhow::ensure!( + self.cluster_timestamp.is_some() || is_key_frame, + "output segment does not begin with a key frame" + ); + if self.cluster_timestamp.is_none() || is_key_frame { + if self.cluster_timestamp.is_some() { + self.writer.write(&MatroskaSpec::Cluster(Master::End))?; + } + self.writer.write_advanced( + &MatroskaSpec::Cluster(Master::Start), + WriteOptions::is_unknown_sized_element(), + )?; + self.writer.write(&MatroskaSpec::Timestamp(timestamp))?; + self.cluster_timestamp = Some(timestamp); + } + + let cluster_timestamp = self.cluster_timestamp.context("output cluster timestamp is missing")?; + let block_timestamp = timestamp + .checked_sub(cluster_timestamp) + .context("output frame timestamp precedes its cluster")?; + let block_timestamp = + i16::try_from(block_timestamp).context("output cluster exceeds block timestamp range")?; + let block = SimpleBlock::new_uncheked(&data, 1, block_timestamp, false, None, false, is_key_frame); + self.writer.write(&MatroskaSpec::from(block))?; + } + + Ok(()) + } + + fn finish(mut self) -> anyhow::Result<()> { + self.encoder.flush()?; + self.write_encoded_frames()?; + if self.cluster_timestamp.is_some() { + self.writer.write(&MatroskaSpec::Cluster(Master::End))?; + } + let event_writer = self.writer.into_inner()?; + send_event(&event_writer.sender, SegmentEvent::End) + .with_context(|| format!("failed to finish segment {}", self.info.sequence)) + } +} + +fn write_header(writer: &mut WebmWriter, width: u32, height: u32) -> anyhow::Result<()> { + writer.write(&MatroskaSpec::Ebml(Master::Full(vec![ + MatroskaSpec::EbmlVersion(1), + MatroskaSpec::EbmlReadVersion(1), + MatroskaSpec::EbmlMaxIdLength(4), + MatroskaSpec::EbmlMaxSizeLength(8), + MatroskaSpec::DocType("webm".to_owned()), + MatroskaSpec::DocTypeVersion(4), + MatroskaSpec::DocTypeReadVersion(2), + ])))?; + writer.write_advanced( + &MatroskaSpec::Segment(Master::Start), + WriteOptions::is_unknown_sized_element(), + )?; + writer.write(&MatroskaSpec::Info(Master::Full(vec![ + MatroskaSpec::TimestampScale(WEBM_TIMESTAMP_SCALE_NS), + MatroskaSpec::MuxingApp("Devolutions Gateway".to_owned()), + MatroskaSpec::WritingApp("Devolutions Gateway".to_owned()), + ])))?; + writer.write(&MatroskaSpec::Tracks(Master::Full(vec![MatroskaSpec::TrackEntry( + Master::Full(vec![ + MatroskaSpec::TrackNumber(1), + MatroskaSpec::TrackUID(1), + MatroskaSpec::TrackType(1), + MatroskaSpec::FlagEnabled(1), + MatroskaSpec::FlagDefault(1), + MatroskaSpec::FlagLacing(0), + MatroskaSpec::CodecID("V_VP8".to_owned()), + MatroskaSpec::Video(Master::Full(vec![ + MatroskaSpec::PixelWidth(u64::from(width)), + MatroskaSpec::PixelHeight(u64::from(height)), + ])), + ]), + )])))?; + Ok(()) +} + +fn send_event(sender: &mpsc::Sender>, event: SegmentEvent) -> anyhow::Result<()> { + sender + .blocking_send(Ok(event)) + .map_err(|_| anyhow::anyhow!("segment event receiver closed")) +} + +struct EventWriter { + sender: mpsc::Sender>, +} + +impl Write for EventWriter { + fn write(&mut self, buffer: &[u8]) -> io::Result { + for chunk in buffer.chunks(OUTPUT_CHUNK_SIZE) { + self.sender + .blocking_send(Ok(SegmentEvent::Data(Bytes::copy_from_slice(chunk)))) + .map_err(|_| io::Error::new(io::ErrorKind::BrokenPipe, "segment event receiver closed"))?; + } + Ok(buffer.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/video-streamer/src/normalizer/tests/mod.rs b/crates/video-streamer/src/normalizer/tests/mod.rs new file mode 100644 index 000000000..edcca2185 --- /dev/null +++ b/crates/video-streamer/src/normalizer/tests/mod.rs @@ -0,0 +1,548 @@ +use std::io::{Cursor, Read, Seek, SeekFrom}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex, mpsc as std_mpsc}; +use std::thread; +use std::time::Duration; + +use super::*; + +#[derive(Default)] +struct ReaderStats { + max_requested: usize, + read_count: usize, + seek_count: usize, + seek_positions: Vec, +} + +struct GrowingReader { + data: Vec, + position: usize, + visible: Arc, + stats: Arc>, +} + +impl Read for GrowingReader { + fn read(&mut self, buffer: &mut [u8]) -> io::Result { + let visible = self.visible.load(Ordering::Acquire).min(self.data.len()); + let available = visible.saturating_sub(self.position); + let read = available.min(buffer.len()); + if read > 0 { + buffer[..read].copy_from_slice(&self.data[self.position..self.position + read]); + self.position += read; + } + let mut stats = self.stats.lock().expect("reader stats lock"); + stats.max_requested = stats.max_requested.max(buffer.len()); + stats.read_count += 1; + Ok(read) + } +} + +impl Seek for GrowingReader { + fn seek(&mut self, position: SeekFrom) -> io::Result { + let next = match position { + SeekFrom::Start(offset) => i64::try_from(offset).map_err(io::Error::other)?, + SeekFrom::Current(offset) => i64::try_from(self.position) + .map_err(io::Error::other)? + .checked_add(offset) + .ok_or_else(|| io::Error::other("seek overflow"))?, + SeekFrom::End(offset) => i64::try_from(self.data.len()) + .map_err(io::Error::other)? + .checked_add(offset) + .ok_or_else(|| io::Error::other("seek overflow"))?, + }; + if next < 0 { + return Err(io::Error::new(io::ErrorKind::InvalidInput, "negative seek")); + } + self.position = usize::try_from(next).map_err(io::Error::other)?; + let mut stats = self.stats.lock().expect("reader stats lock"); + stats.seek_count += 1; + stats + .seek_positions + .push(u64::try_from(self.position).map_err(io::Error::other)?); + u64::try_from(self.position).map_err(io::Error::other) + } +} + +fn video_clip_bytes(key_frames: &[bool], block_group: bool) -> Vec { + video_clip_bytes_with_group_sizing(key_frames, block_group, false) +} + +fn video_clip_bytes_with_group_sizing(key_frames: &[bool], block_group: bool, unknown_block_group: bool) -> Vec { + use webm_iterable::matroska_spec::Block; + + let mut writer = WebmWriter::new(Vec::new()); + writer + .write(&MatroskaSpec::Ebml(Master::Full(vec![ + MatroskaSpec::EbmlVersion(1), + MatroskaSpec::EbmlReadVersion(1), + MatroskaSpec::EbmlMaxIdLength(4), + MatroskaSpec::EbmlMaxSizeLength(8), + MatroskaSpec::DocType("webm".to_owned()), + MatroskaSpec::DocTypeVersion(4), + MatroskaSpec::DocTypeReadVersion(2), + ]))) + .expect("write EBML header"); + writer + .write_advanced( + &MatroskaSpec::Segment(Master::Start), + WriteOptions::is_unknown_sized_element(), + ) + .expect("write segment start"); + writer + .write(&MatroskaSpec::Info(Master::Full(vec![MatroskaSpec::TimestampScale( + WEBM_TIMESTAMP_SCALE_NS, + )]))) + .expect("write info"); + writer + .write(&MatroskaSpec::Tracks(Master::Full(vec![MatroskaSpec::TrackEntry( + Master::Full(vec![ + MatroskaSpec::TrackNumber(1), + MatroskaSpec::TrackType(1), + MatroskaSpec::CodecID("V_VP8".to_owned()), + ]), + )]))) + .expect("write video track"); + writer + .write_advanced( + &MatroskaSpec::Cluster(Master::Start), + WriteOptions::is_unknown_sized_element(), + ) + .expect("write cluster start"); + writer + .write(&MatroskaSpec::Timestamp(0)) + .expect("write cluster timestamp"); + + for (index, &key_frame) in key_frames.iter().enumerate() { + let frame = if key_frame { [0] } else { [1] }; + let timestamp = i16::try_from(index * 30).expect("test timestamp fits"); + if block_group { + if unknown_block_group { + writer + .write_advanced( + &MatroskaSpec::BlockGroup(Master::Start), + WriteOptions::is_unknown_sized_element(), + ) + .expect("write unknown-sized block group start"); + } else { + writer + .write(&MatroskaSpec::BlockGroup(Master::Start)) + .expect("write block group start"); + } + writer + .write(&MatroskaSpec::from(Block::new_uncheked( + 1, timestamp, false, None, &frame, + ))) + .expect("write block"); + writer + .write(&MatroskaSpec::BlockGroup(Master::End)) + .expect("write block group end"); + } else { + writer + .write(&MatroskaSpec::from(SimpleBlock::new_uncheked( + &frame, 1, timestamp, false, None, false, key_frame, + ))) + .expect("write simple block"); + } + } + + if block_group { + writer + .write(&MatroskaSpec::Timestamp(1)) + .expect("write block group terminator"); + } + + writer.into_inner().expect("finish video clip bytes") +} + +fn truncate_inside_block_payload(data: &[u8], block_index: usize) -> Vec { + let mut decoder = TagDecoder::new(&[]); + let mut input = BytesMut::from(data); + let mut current_index = 0; + while let Some(positioned) = decoder.decode(&mut input).expect("decode complete fixture") { + if matches!(positioned.tag, MatroskaSpec::Block(_)) { + if current_index == block_index { + let block_end = decoder.position(); + let mut truncated = data.to_vec(); + truncated.truncate(block_end.checked_sub(1).expect("block payload is nonempty")); + return truncated; + } + current_index += 1; + } + } + panic!("fixture does not contain requested block"); +} + +fn truncate_before_last_timestamp(data: &[u8]) -> Vec { + let mut decoder = TagDecoder::new(&[]); + let mut input = BytesMut::from(data); + let mut last_timestamp_start = None; + while let Some(positioned) = decoder.decode(&mut input).expect("decode complete fixture") { + if matches!(positioned.tag, MatroskaSpec::Timestamp(_)) { + last_timestamp_start = Some(positioned.offset); + } + } + let timestamp_start = last_timestamp_start.expect("fixture has a timestamp terminator"); + data[..timestamp_start].to_vec() +} + +fn live_edge_normalizer(reader: R) -> (ClipNormalizer, mpsc::Receiver>) +where + R: Read + Seek + Send + 'static, +{ + let (sender, receiver) = mpsc::channel(8); + let normalizer = ClipNormalizer::new( + 0, + StartAt::LiveEdge, + RecordingClip::new(reader), + sender, + SessionConfig { + encoder_threads: 1, + adaptive_frame_skip: false, + }, + 0, + ) + .expect("create live-edge normalizer"); + (normalizer, receiver) +} + +#[test] +fn adaptive_frame_skip_is_bounded() { + let (mut normalizer, _receiver) = live_edge_normalizer(Cursor::new(Vec::new())); + normalizer.config.adaptive_frame_skip = true; + normalizer.first_frame_timestamp = Some(0); + normalizer.processing_time = Duration::from_millis(100); + + assert!(normalizer.should_skip_encode(50)); + normalizer.frames_since_last_encode = 1; + assert!(!normalizer.should_skip_encode(50)); + + normalizer.frames_since_last_encode = 0; + normalizer.config.adaptive_frame_skip = false; + assert!(!normalizer.should_skip_encode(50)); +} + +fn empty_clip_bytes() -> Vec { + let mut writer = WebmWriter::new(Vec::new()); + writer + .write(&MatroskaSpec::Ebml(Master::Full(vec![ + MatroskaSpec::EbmlVersion(1), + MatroskaSpec::EbmlReadVersion(1), + MatroskaSpec::EbmlMaxIdLength(4), + MatroskaSpec::EbmlMaxSizeLength(8), + MatroskaSpec::DocType("webm".to_owned()), + MatroskaSpec::DocTypeVersion(4), + MatroskaSpec::DocTypeReadVersion(2), + ]))) + .expect("write EBML header"); + writer + .write_advanced( + &MatroskaSpec::Segment(Master::Start), + WriteOptions::is_unknown_sized_element(), + ) + .expect("write segment start"); + writer + .write_advanced( + &MatroskaSpec::Cluster(Master::Start), + WriteOptions::is_unknown_sized_element(), + ) + .expect("write cluster start"); + writer + .write(&MatroskaSpec::Timestamp(0)) + .expect("write cluster timestamp"); + writer.into_inner().expect("finish clip bytes") +} + +mod replay; + +#[test] +fn resolution_change_starts_the_next_output_segment() { + let first_dimensions = Dimensions { + width: 640, + height: 480, + }; + let second_dimensions = Dimensions { + width: 1280, + height: 720, + }; + + assert_eq!( + next_segment_info(None, first_dimensions, 0), + Some(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + }) + ); + assert_eq!(next_segment_info(Some(first_dimensions), first_dimensions, 1), None); + assert_eq!( + next_segment_info(Some(first_dimensions), second_dimensions, 1), + Some(SegmentInfo { + sequence: 1, + width: 1280, + height: 720, + }) + ); +} + +#[test] +fn scanner_uses_bounded_reads_and_constant_replay_metadata() { + let data = video_clip_bytes(&[true, false, false, true, false], false); + let visible = Arc::new(AtomicUsize::new(data.len())); + let stats = Arc::new(Mutex::new(ReaderStats::default())); + let reader = GrowingReader { + data, + position: 0, + visible, + stats: Arc::clone(&stats), + }; + let (mut normalizer, _receiver) = live_edge_normalizer(reader); + + normalizer.scan_available().expect("scan available history"); + + let stats = stats.lock().expect("reader stats lock"); + assert!(stats.max_requested <= INPUT_CHUNK_SIZE); + assert_eq!(stats.seek_count, 1); + assert!(normalizer.input.is_empty()); + assert!(normalizer.replay_point.is_some()); + assert!(normalizer.complete_boundary > normalizer.replay_point.expect("replay point").block_offset); +} + +#[test] +fn incomplete_simple_block_continues_once_when_reader_grows() { + let data = video_clip_bytes(&[true], false); + let visible = Arc::new(AtomicUsize::new(data.len() - 1)); + let reader = GrowingReader { + data: data.clone(), + position: 0, + visible: Arc::clone(&visible), + stats: Arc::new(Mutex::new(ReaderStats::default())), + }; + let (mut normalizer, _receiver) = live_edge_normalizer(reader); + + normalizer.scan_available().expect("scan partial simple block"); + assert!(normalizer.replay_point.is_none()); + + visible.store(data.len(), Ordering::Release); + normalizer.scan_available().expect("continue simple block"); + assert!(normalizer.replay_point.is_some()); + assert_eq!( + normalizer.complete_boundary, + u64::try_from(data.len()).expect("fixture length fits") + ); +} + +#[test] +fn incomplete_block_group_does_not_become_a_replay_boundary() { + let data = video_clip_bytes(&[true], true); + let partial = truncate_inside_block_payload(&data, 0); + let visible = Arc::new(AtomicUsize::new(partial.len())); + let reader = GrowingReader { + data: data.clone(), + position: 0, + visible: Arc::clone(&visible), + stats: Arc::new(Mutex::new(ReaderStats::default())), + }; + let (mut normalizer, _receiver) = live_edge_normalizer(reader); + + normalizer.scan_available().expect("scan partial block group"); + assert!(normalizer.replay_point.is_none()); + assert_eq!(normalizer.complete_boundary, 0); + + visible.store(data.len(), Ordering::Release); + normalizer.scan_available().expect("continue block group"); + assert!(normalizer.replay_point.is_some()); + assert!(normalizer.complete_boundary > 0); +} + +#[test] +fn oversized_element_is_rejected_before_unbounded_capture() { + let oversized = vec![0xa3, 0x14, 0x00, 0x00, 0x01]; + let (mut normalizer, _receiver) = live_edge_normalizer(Cursor::new(oversized)); + + assert!(normalizer.scan_available().is_err()); +} + +#[test] +fn output_data_events_are_bounded() { + let (sender, mut receiver) = mpsc::channel(4); + let mut writer = EventWriter { sender }; + let data = vec![0; OUTPUT_CHUNK_SIZE * 2 + 1]; + + assert_eq!(writer.write(&data).expect("write output data"), data.len()); + for expected_len in [OUTPUT_CHUNK_SIZE, OUTPUT_CHUNK_SIZE, 1] { + let event = receiver + .blocking_recv() + .expect("receive output data") + .expect("output event"); + let SegmentEvent::Data(data) = event else { + panic!("unexpected output event"); + }; + assert_eq!(data.len(), expected_len); + } +} + +#[test] +fn output_prefetch_is_bounded_and_refills_in_order() { + let (sender, mut receiver) = mpsc::channel(OUTPUT_CHANNEL_CAPACITY); + let mut writer = EventWriter { sender }; + let (partial_ready_sender, partial_ready_receiver) = std_mpsc::channel(); + let (prefetch_release_sender, prefetch_release_receiver) = std_mpsc::channel(); + let (prefetch_ready_sender, prefetch_ready_receiver) = std_mpsc::channel(); + let (refill_started_sender, refill_started_receiver) = std_mpsc::channel(); + let (finished_sender, finished_receiver) = std_mpsc::channel(); + let producer = thread::spawn(move || { + writer.write_all(b"partial").expect("write partial output"); + partial_ready_sender.send(()).expect("signal partial output"); + prefetch_release_receiver.recv().expect("release prefetch"); + + let mut prefetched = Vec::with_capacity(OUTPUT_CHANNEL_CAPACITY * OUTPUT_CHUNK_SIZE); + for marker in 1..=OUTPUT_CHANNEL_CAPACITY { + prefetched.extend(std::iter::repeat_n( + u8::try_from(marker).expect("marker fits in u8"), + OUTPUT_CHUNK_SIZE, + )); + } + writer.write_all(&prefetched).expect("write prefetched output"); + prefetch_ready_sender.send(()).expect("signal full prefetch"); + + refill_started_sender.send(()).expect("signal refill attempt"); + let mut refill = Vec::with_capacity(2 * OUTPUT_CHUNK_SIZE); + for marker in OUTPUT_CHANNEL_CAPACITY + 1..=OUTPUT_CHANNEL_CAPACITY + 2 { + refill.extend(std::iter::repeat_n( + u8::try_from(marker).expect("marker fits in u8"), + OUTPUT_CHUNK_SIZE, + )); + } + writer.write_all(&refill).expect("write refill output"); + finished_sender.send(()).expect("signal producer completion"); + }); + + partial_ready_receiver + .recv_timeout(Duration::from_secs(1)) + .expect("partial output was not ready"); + assert_eq!(receiver.len(), 1); + let partial = receiver + .try_recv() + .expect("receive partial output") + .expect("partial output event"); + assert_eq!(partial, SegmentEvent::Data(Bytes::from_static(b"partial"))); + + prefetch_release_sender.send(()).expect("release prefetch"); + prefetch_ready_receiver + .recv_timeout(Duration::from_secs(1)) + .expect("prefetch did not fill"); + assert_eq!(receiver.len(), OUTPUT_CHANNEL_CAPACITY); + assert_eq!(receiver.len() * OUTPUT_CHUNK_SIZE, 256 * 1024); + refill_started_receiver + .recv_timeout(Duration::from_secs(1)) + .expect("refill did not start"); + assert!(matches!( + finished_receiver.recv_timeout(Duration::from_millis(25)), + Err(std_mpsc::RecvTimeoutError::Timeout) + )); + + for expected_marker in 1..=OUTPUT_CHANNEL_CAPACITY { + let event = receiver + .try_recv() + .expect("receive prefetched output") + .expect("prefetched output event"); + let SegmentEvent::Data(data) = event else { + panic!("unexpected output event"); + }; + assert_eq!(data.len(), OUTPUT_CHUNK_SIZE); + let expected_marker = u8::try_from(expected_marker).expect("marker fits in u8"); + assert!(data.iter().all(|&byte| byte == expected_marker)); + } + finished_receiver + .recv_timeout(Duration::from_secs(1)) + .expect("producer did not finish after consumption"); + for expected_marker in OUTPUT_CHANNEL_CAPACITY + 1..=OUTPUT_CHANNEL_CAPACITY + 2 { + let event = receiver + .try_recv() + .expect("receive refilled output") + .expect("refilled output event"); + let SegmentEvent::Data(data) = event else { + panic!("unexpected output event"); + }; + assert_eq!(data.len(), OUTPUT_CHUNK_SIZE); + let expected_marker = u8::try_from(expected_marker).expect("marker fits in u8"); + assert!(data.iter().all(|&byte| byte == expected_marker)); + } + producer.join().expect("producer thread panicked"); +} + +#[test] +fn truncated_clip_tail_does_not_abort_the_following_clip() { + let mut truncated = empty_clip_bytes(); + truncated.extend_from_slice(&[0xa3, 0x84, 0x81, 0x00]); + let complete = empty_clip_bytes(); + let events = vec![ + RecordingEvent::ClipStarted { + sequence: 0, + start_at: StartAt::Beginning, + clip: RecordingClip::new(Cursor::new(truncated)), + }, + RecordingEvent::CaughtUp, + RecordingEvent::ClipEnded, + RecordingEvent::ClipStarted { + sequence: 1, + start_at: StartAt::Beginning, + clip: RecordingClip::new(Cursor::new(complete)), + }, + RecordingEvent::CaughtUp, + RecordingEvent::ClipEnded, + RecordingEvent::SessionEnded, + ]; + let (input_sender, input_receiver) = mpsc::channel(events.len()); + for event in events { + input_sender.blocking_send(Ok(event)).expect("queue recording event"); + } + drop(input_sender); + let (output_sender, mut output_receiver) = mpsc::channel(1); + + normalize_events( + input_receiver, + output_sender, + SessionConfig { + encoder_threads: 1, + adaptive_frame_skip: false, + }, + ) + .expect("normalize reconnecting clips"); + assert!(output_receiver.blocking_recv().is_none()); +} + +#[test] +fn corruption_before_an_incomplete_tail_still_fails() { + let mut corrupted = empty_clip_bytes(); + corrupted.extend_from_slice(&[0xff, 0x80]); + corrupted.extend_from_slice(&[0xa3, 0x84, 0x81, 0x00]); + let events = vec![ + RecordingEvent::ClipStarted { + sequence: 0, + start_at: StartAt::Beginning, + clip: RecordingClip::new(Cursor::new(corrupted)), + }, + RecordingEvent::CaughtUp, + RecordingEvent::ClipEnded, + RecordingEvent::SessionEnded, + ]; + let (input_sender, input_receiver) = mpsc::channel(events.len()); + for event in events { + input_sender.blocking_send(Ok(event)).expect("queue recording event"); + } + drop(input_sender); + let (output_sender, _output_receiver) = mpsc::channel(1); + + let error = normalize_events( + input_receiver, + output_sender, + SessionConfig { + encoder_threads: 1, + adaptive_frame_skip: false, + }, + ) + .expect_err("corruption before the incomplete tail must fail"); + + assert!(format!("{error:#}").contains("corrupted"), "{error:#}"); +} diff --git a/crates/video-streamer/src/normalizer/tests/replay.rs b/crates/video-streamer/src/normalizer/tests/replay.rs new file mode 100644 index 000000000..179450056 --- /dev/null +++ b/crates/video-streamer/src/normalizer/tests/replay.rs @@ -0,0 +1,293 @@ +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; + +use tokio::sync::mpsc; +use webm_iterable::matroska_spec::{Master, MatroskaSpec, SimpleBlock}; +use webm_iterable::{WebmWriter, WriteOptions}; + +use super::*; + +#[test] +fn replay_uses_latest_gop_and_restores_the_same_reader() { + let data = video_clip_bytes(&[true, false, false, true, false], false); + let visible = Arc::new(AtomicUsize::new(data.len())); + let stats = Arc::new(Mutex::new(ReaderStats::default())); + let reader = GrowingReader { + data, + position: 0, + visible, + stats: Arc::clone(&stats), + }; + let (mut normalizer, _receiver) = live_edge_normalizer(reader); + + normalizer.scan_available().expect("scan available history"); + let original_head = normalizer.reader_head; + let replay_point = normalizer.replay_point.expect("latest replay point"); + let mut frames = Vec::new(); + let mut process_frame = |_: &mut ClipNormalizer, frame: PendingFrame| { + frames.push((frame.timestamp, frame.key_frame)); + Ok(()) + }; + + normalizer + .caught_up_with(&mut process_frame) + .expect("replay latest GOP"); + + assert_eq!(frames, vec![(90, true), (120, false)]); + assert_eq!(normalizer.reader_head, original_head); + + let stats = stats.lock().expect("reader stats lock"); + assert_eq!(stats.seek_positions.first(), Some(&0)); + assert_eq!(stats.seek_positions.get(1), Some(&replay_point.block_offset)); + assert_eq!(stats.seek_positions.last(), Some(&original_head)); +} + +#[test] +fn replay_error_restores_the_original_parser_and_reader_state() { + let data = video_clip_bytes(&[true, false, true], false); + let visible = Arc::new(AtomicUsize::new(data.len())); + let stats = Arc::new(Mutex::new(ReaderStats::default())); + let reader = GrowingReader { + data, + position: 0, + visible, + stats: Arc::clone(&stats), + }; + let (mut normalizer, _receiver) = live_edge_normalizer(reader); + + normalizer.scan_available().expect("scan available history"); + let original_head = normalizer.reader_head; + let original_decoder_position = normalizer.decoder.position(); + let original_input = normalizer.input.clone(); + let mut process_frame = |_: &mut ClipNormalizer, _: PendingFrame| Err(anyhow::anyhow!("replay callback failed")); + + let error = normalizer + .caught_up_with(&mut process_frame) + .expect_err("replay callback failure"); + + assert!(format!("{error:#}").contains("replay callback failed")); + assert_eq!(normalizer.reader_head, original_head); + assert_eq!(normalizer.decoder.position(), original_decoder_position); + assert_eq!(normalizer.input, original_input); + + let stats = stats.lock().expect("reader stats lock"); + assert_eq!(stats.seek_positions.last(), Some(&original_head)); +} + +#[test] +fn known_size_group_replay_excludes_a_partial_next_group_until_growth() { + let data = video_clip_bytes(&[true, true], true); + let partial = truncate_inside_block_payload(&data, 1); + let visible = Arc::new(AtomicUsize::new(partial.len())); + let reader = GrowingReader { + data: data.clone(), + position: 0, + visible: Arc::clone(&visible), + stats: Arc::new(Mutex::new(ReaderStats::default())), + }; + let (mut normalizer, _receiver) = live_edge_normalizer(reader); + let mut frames = Vec::new(); + + normalizer.scan_available().expect("scan partial group"); + assert!(normalizer.pending_block_group.is_some()); + { + let mut process_frame = |_: &mut ClipNormalizer, frame: PendingFrame| { + frames.push((frame.timestamp, frame.key_frame)); + Ok(()) + }; + normalizer + .caught_up_with(&mut process_frame) + .expect("replay complete group"); + } + assert_eq!(frames, vec![(0, true)]); + + visible.store(data.len(), Ordering::Release); + let mut process_frame = |_: &mut ClipNormalizer, frame: PendingFrame| { + frames.push((frame.timestamp, frame.key_frame)); + Ok(()) + }; + normalizer + .scan_available_with(&mut process_frame) + .expect("consume grown group"); + assert_eq!(frames, vec![(0, true), (30, true)]); +} + +#[test] +fn unknown_size_group_replay_uses_the_following_sibling_as_its_boundary() { + let data = video_clip_bytes_with_group_sizing(&[true, true], true, true); + let partial = truncate_before_last_timestamp(&data); + let visible = Arc::new(AtomicUsize::new(partial.len())); + let reader = GrowingReader { + data: data.clone(), + position: 0, + visible: Arc::clone(&visible), + stats: Arc::new(Mutex::new(ReaderStats::default())), + }; + let (mut normalizer, _receiver) = live_edge_normalizer(reader); + let mut frames = Vec::new(); + + normalizer.scan_available().expect("scan unknown-sized group"); + assert!(normalizer.pending_block_group.is_some()); + { + let mut process_frame = |_: &mut ClipNormalizer, frame: PendingFrame| { + frames.push((frame.timestamp, frame.key_frame)); + Ok(()) + }; + normalizer + .caught_up_with(&mut process_frame) + .expect("replay through the complete boundary"); + } + assert_eq!(frames, vec![(0, true)]); + + visible.store(data.len(), Ordering::Release); + let mut process_frame = |_: &mut ClipNormalizer, frame: PendingFrame| { + frames.push((frame.timestamp, frame.key_frame)); + Ok(()) + }; + normalizer + .scan_available_with(&mut process_frame) + .expect("consume following sibling"); + assert_eq!(frames, vec![(0, true), (30, true)]); +} + +#[test] +fn replay_stops_at_a_completed_frame_inside_a_known_cluster() { + let data = cross_cluster_replay_bytes(); + let visible = Arc::new(AtomicUsize::new(data.len() - 1)); + let reader = GrowingReader { + data, + position: 0, + visible, + stats: Arc::new(Mutex::new(ReaderStats::default())), + }; + let (mut normalizer, _receiver) = live_edge_normalizer(reader); + let mut frames = Vec::new(); + let mut process_frame = |_: &mut ClipNormalizer, frame: PendingFrame| { + frames.push((frame.timestamp, frame.key_frame)); + Ok(()) + }; + + normalizer.scan_available().expect("scan cross-cluster history"); + normalizer + .caught_up_with(&mut process_frame) + .expect("replay completed frames without finalizing the document"); + + assert_eq!(frames, vec![(0, true), (30, false)]); +} + +#[test] +fn closed_output_stops_scan_at_a_frame_checkpoint() { + let data = video_clip_bytes(&[true, false, false], false); + let visible = Arc::new(AtomicUsize::new(data.len())); + let stats = Arc::new(Mutex::new(ReaderStats::default())); + let reader = GrowingReader { + data, + position: 0, + visible, + stats: Arc::clone(&stats), + }; + let (sender, mut receiver) = mpsc::channel(8); + let mut normalizer = ClipNormalizer::new( + 0, + StartAt::Beginning, + RecordingClip::new(reader), + sender, + SessionConfig { + encoder_threads: 1, + adaptive_frame_skip: false, + }, + 0, + ) + .expect("create beginning normalizer"); + let mut process_frame = |_: &mut ClipNormalizer, _: PendingFrame| { + receiver.close(); + Ok(()) + }; + + normalizer + .scan_available_with(&mut process_frame) + .expect("stop scan after receiver closes"); + + let stats = stats.lock().expect("reader stats lock"); + assert_eq!(stats.read_count, 1); +} + +#[test] +fn zero_timestamp_scale_remains_accepted() { + let (mut normalizer, _receiver) = live_edge_normalizer(Cursor::new(empty_clip_bytes())); + + normalizer + .handle_tag(PositionedTag { + tag: MatroskaSpec::TimestampScale(0), + offset: 0, + }) + .expect("accept timestamp scale"); + + assert_eq!(normalizer.timestamp_scale_ns, 0); +} + +fn cross_cluster_replay_bytes() -> Vec { + let mut writer = WebmWriter::new(Vec::new()); + writer + .write(&MatroskaSpec::Ebml(Master::Full(vec![ + MatroskaSpec::EbmlVersion(1), + MatroskaSpec::EbmlReadVersion(1), + MatroskaSpec::EbmlMaxIdLength(4), + MatroskaSpec::EbmlMaxSizeLength(8), + MatroskaSpec::DocType("webm".to_owned()), + MatroskaSpec::DocTypeVersion(4), + MatroskaSpec::DocTypeReadVersion(2), + ]))) + .expect("write EBML header"); + writer + .write_advanced( + &MatroskaSpec::Segment(Master::Start), + WriteOptions::is_unknown_sized_element(), + ) + .expect("write segment start"); + writer + .write(&MatroskaSpec::Info(Master::Full(vec![MatroskaSpec::TimestampScale( + WEBM_TIMESTAMP_SCALE_NS, + )]))) + .expect("write info"); + writer + .write(&MatroskaSpec::Tracks(Master::Full(vec![MatroskaSpec::TrackEntry( + Master::Full(vec![ + MatroskaSpec::TrackNumber(1), + MatroskaSpec::TrackType(1), + MatroskaSpec::CodecID("V_VP8".to_owned()), + ]), + )]))) + .expect("write video track"); + writer + .write_advanced( + &MatroskaSpec::Cluster(Master::Start), + WriteOptions::is_unknown_sized_element(), + ) + .expect("write first cluster start"); + writer + .write(&MatroskaSpec::Timestamp(0)) + .expect("write first cluster timestamp"); + writer + .write(&MatroskaSpec::from(SimpleBlock::new_uncheked( + &[0], + 1, + 0, + false, + None, + false, + true, + ))) + .expect("write first frame"); + writer + .write(&MatroskaSpec::Cluster(Master::End)) + .expect("write first cluster end"); + writer + .write(&MatroskaSpec::Cluster(Master::Full(vec![ + MatroskaSpec::Timestamp(30), + MatroskaSpec::from(SimpleBlock::new_uncheked(&[1], 1, 0, false, None, false, false)), + MatroskaSpec::from(SimpleBlock::new_uncheked(&[1], 1, 30, false, None, false, false)), + ]))) + .expect("write known-size second cluster"); + writer.into_inner().expect("finish cross-cluster fixture") +} diff --git a/crates/video-streamer/src/protocol/message.rs b/crates/video-streamer/src/protocol/message.rs new file mode 100644 index 000000000..242bafa44 --- /dev/null +++ b/crates/video-streamer/src/protocol/message.rs @@ -0,0 +1,75 @@ +use bytes::{BufMut as _, Bytes, BytesMut}; + +const VP8_METADATA_PAYLOAD: &[u8] = b"{\"codec\":\"vp8\"}"; + +#[derive(Debug, Eq, PartialEq)] +pub(super) enum ServerMessage { + Chunk(Bytes), + Metadata, + SegmentStarted, + Error(UserFriendlyError), + StreamEnded, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(super) enum ClientMessage { + Start, + Pull, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub(super) enum UserFriendlyError { + UnexpectedError, +} + +impl UserFriendlyError { + fn as_str(&self) -> &'static str { + match self { + Self::UnexpectedError => "UnexpectedError", + } + } +} + +pub(super) fn decode_client_message(message: &[u8]) -> anyhow::Result { + match message { + [0] => Ok(ClientMessage::Start), + [1] => Ok(ClientMessage::Pull), + _ => anyhow::bail!("invalid client message"), + } +} + +pub(super) fn response_kind(message: &ServerMessage) -> &'static str { + match message { + ServerMessage::Chunk(_) => "chunk", + ServerMessage::Metadata => "metadata", + ServerMessage::SegmentStarted => "segment-started", + ServerMessage::Error(_) => "error", + ServerMessage::StreamEnded => "stream-ended", + } +} + +pub(super) fn encode_server_message(message: ServerMessage) -> Bytes { + let mut encoded = BytesMut::new(); + match message { + ServerMessage::Chunk(chunk) => { + encoded.reserve(1 + chunk.len()); + encoded.put_u8(0); + encoded.put(chunk); + } + ServerMessage::Metadata => { + encoded.put_u8(1); + encoded.put(VP8_METADATA_PAYLOAD); + } + ServerMessage::Error(error) => { + encoded.put_u8(2); + let json = format!("{{\"error\":\"{}\"}}", error.as_str()); + encoded.put(json.as_bytes()); + } + ServerMessage::StreamEnded => encoded.put_u8(3), + ServerMessage::SegmentStarted => { + encoded.put_u8(4); + encoded.put(VP8_METADATA_PAYLOAD); + } + } + encoded.freeze() +} diff --git a/crates/video-streamer/src/protocol/mod.rs b/crates/video-streamer/src/protocol/mod.rs new file mode 100644 index 000000000..794f2a775 --- /dev/null +++ b/crates/video-streamer/src/protocol/mod.rs @@ -0,0 +1,113 @@ +use std::error::Error; + +use bytes::Bytes; +use futures_util::{Sink, Stream}; + +use crate::normalizer::NormalizedSession; +use crate::session::{RecordingSource, SessionConfig}; + +mod message; +mod segments; +mod transport; + +use message::{ClientMessage, ServerMessage, response_kind}; +use segments::SessionSegments; +use transport::{CodecTransport, ReceiveError, SessionTransport}; + +pub(crate) async fn stream_segments(transport: T, source: S, config: SessionConfig) -> anyhow::Result<()> +where + S: RecordingSource, + T: Stream> + Sink + Unpin, + E: Error + Send + Sync + 'static, +{ + let mut transport = SessionTransport::new(CodecTransport::new(transport)); + let Some(_) = receive_expected_request(&mut transport, ClientMessage::Start).await? else { + return Ok(()); + }; + + let source_stream = match source.start().await { + Ok(source_stream) => source_stream, + Err(error) => { + transport.reject().await; + return Err(error); + } + }; + let mut segments = SessionSegments::new(crate::normalizer::normalize(source_stream, config)); + let stream_result = run_started_session(&mut transport, &mut segments).await; + let shutdown_result = segments.into_inner().shutdown().await; + + match (stream_result, shutdown_result) { + (Err(error), _) => Err(error), + (Ok(()), Err(error)) => Err(error), + (Ok(()), Ok(())) => Ok(()), + } +} + +async fn receive_expected_request( + transport: &mut SessionTransport, + expected: ClientMessage, +) -> anyhow::Result> +where + T: Stream>> + Sink + Unpin, + E: Error + Send + Sync + 'static, +{ + let Some(incoming) = transport.recv().await else { + return Ok(None); + }; + let message = match incoming { + Ok(message) => message, + Err(ReceiveError::Transport(error)) => { + return Err(anyhow::Error::new(error).context("read client stream message")); + } + Err(ReceiveError::Decode(error)) => { + debug!(error = %error, "Rejected undecodable client request"); + transport.reject().await; + return Err(error.context("decode client request")); + } + }; + + if message != expected { + debug!(expected = ?expected, got = ?message, "Rejected client request in wrong state"); + transport.reject().await; + anyhow::bail!("invalid client stream state"); + } + + Ok(Some(message)) +} + +async fn run_started_session( + transport: &mut SessionTransport, + segments: &mut SessionSegments, +) -> anyhow::Result<()> +where + T: Stream>> + Sink + Unpin, + E: Error + Send + Sync + 'static, +{ + debug!("Serving Start request"); + transport.send(ServerMessage::Metadata).await?; + + loop { + let Some(_) = receive_expected_request(transport, ClientMessage::Pull).await? else { + return Ok(()); + }; + debug!("Serving Pull request"); + let response = match segments.next().await { + Ok(response) => response, + Err(error) => { + debug!(error = %error, "Request failed while waiting"); + transport.reject().await; + return Err(error); + } + }; + + debug!(response = ?response_kind(&response), "Sending server response"); + let ended = response == ServerMessage::StreamEnded; + transport.send(response).await?; + if ended { + return Ok(()); + } + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/video-streamer/src/protocol/segments.rs b/crates/video-streamer/src/protocol/segments.rs new file mode 100644 index 000000000..c482caf13 --- /dev/null +++ b/crates/video-streamer/src/protocol/segments.rs @@ -0,0 +1,89 @@ +use std::pin::Pin; + +use anyhow::Context as _; +use futures_util::{Stream, StreamExt as _}; + +use super::message::ServerMessage; +use crate::normalizer::SegmentEvent; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum SegmentState { + AwaitingBegin { next_sequence: u64 }, + Streaming { next_sequence: u64 }, +} + +pub(super) struct SessionSegments { + inner: Pin>, + state: SegmentState, +} + +impl SessionSegments +where + S: Stream>, +{ + pub(super) fn new(inner: S) -> Self { + Self { + inner: Box::pin(inner), + state: SegmentState::AwaitingBegin { next_sequence: 0 }, + } + } + + pub(super) fn into_inner(self) -> S + where + S: Unpin, + { + *Pin::into_inner(self.inner) + } + + pub(super) async fn next(&mut self) -> anyhow::Result { + loop { + let Some(event) = self.inner.as_mut().next().await else { + anyhow::ensure!( + matches!(self.state, SegmentState::AwaitingBegin { .. }), + "segment stream ended inside a segment" + ); + return Ok(ServerMessage::StreamEnded); + }; + + match event? { + SegmentEvent::Begin(info) => { + let SegmentState::AwaitingBegin { next_sequence } = self.state else { + anyhow::bail!("segment began before the previous segment ended"); + }; + anyhow::ensure!( + info.sequence == next_sequence, + "segment sequence is not contiguous: expected {next_sequence}, got {}", + info.sequence + ); + self.state = SegmentState::Streaming { + next_sequence: next_sequence.checked_add(1).context("segment sequence overflow")?, + }; + debug!( + sequence = info.sequence, + width = info.width, + height = info.height, + "Segment begin" + ); + if info.sequence > 0 { + return Ok(ServerMessage::SegmentStarted); + } + } + SegmentEvent::Data(data) => { + anyhow::ensure!( + matches!(self.state, SegmentState::Streaming { .. }), + "segment data arrived outside a segment" + ); + debug!(bytes = data.len(), "Segment data"); + return Ok(ServerMessage::Chunk(data)); + } + SegmentEvent::End => { + let SegmentState::Streaming { next_sequence } = self.state else { + anyhow::bail!("segment ended outside a segment"); + }; + self.state = SegmentState::AwaitingBegin { next_sequence }; + debug!("Segment end"); + } + } + } + } +} diff --git a/crates/video-streamer/src/protocol/tests/mod.rs b/crates/video-streamer/src/protocol/tests/mod.rs new file mode 100644 index 000000000..846c21d17 --- /dev/null +++ b/crates/video-streamer/src/protocol/tests/mod.rs @@ -0,0 +1,928 @@ +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::task::{Context, Poll}; +use std::time::Duration; + +use bytes::Bytes; +use futures_util::{Sink, Stream, StreamExt as _, stream}; +use tokio::sync::{mpsc, oneshot}; + +use super::message::{ClientMessage, ServerMessage, decode_client_message, encode_server_message}; +use super::segments::SessionSegments; +use super::transport::CodecTransport; +use super::*; +use crate::normalizer::{SegmentEvent, SegmentInfo}; +use crate::session::{RecordingEvent, RecordingSource}; + +struct ChannelTransport { + incoming: mpsc::UnboundedReceiver>, + outgoing: mpsc::UnboundedSender, +} + +struct ClientSender(mpsc::UnboundedSender>); + +impl ClientSender { + fn send(&self, message: Bytes) -> Result<(), mpsc::error::SendError>> { + self.0.send(Ok(message)) + } + + fn send_error(&self, error: std::io::Error) -> Result<(), mpsc::error::SendError>> { + self.0.send(Err(error)) + } +} + +impl Stream for ChannelTransport { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.incoming.poll_recv(cx) + } +} + +impl Sink for ChannelTransport { + type Error = std::io::Error; + + fn poll_ready(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn start_send(self: Pin<&mut Self>, message: Bytes) -> Result<(), Self::Error> { + self.outgoing + .send(message) + .map_err(|_| std::io::Error::new(std::io::ErrorKind::BrokenPipe, "test receiver closed")) + } + + fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } +} + +fn channel_transport() -> (ChannelTransport, ClientSender, mpsc::UnboundedReceiver) { + let (client_sender, incoming) = mpsc::unbounded_channel(); + let (outgoing, client_receiver) = mpsc::unbounded_channel(); + ( + ChannelTransport { incoming, outgoing }, + ClientSender(client_sender), + client_receiver, + ) +} + +async fn receive_response(receiver: &mut mpsc::UnboundedReceiver) -> Bytes { + tokio::time::timeout(Duration::from_secs(1), receiver.recv()) + .await + .expect("timed out waiting for server response") + .expect("server response channel closed") +} + +struct TestRecordingSource(F); + +fn recording_source(start: F) -> TestRecordingSource { + TestRecordingSource(start) +} + +impl RecordingSource for TestRecordingSource +where + F: FnOnce() -> Fut + Send + 'static, + Fut: Future> + Send + 'static, + S: Stream> + Send + 'static, +{ + type Stream = S; + type Start = Fut; + + fn start(self) -> Self::Start { + self.0() + } +} + +fn segment_source(events: impl IntoIterator>) -> Vec> { + events.into_iter().collect() +} + +async fn stream_segment_source(transport: T, source: S) -> anyhow::Result<()> +where + T: Stream> + Sink + Unpin, + E: Error + Send + Sync + 'static, + S: Stream> + Send + 'static, +{ + let mut transport = SessionTransport::new(CodecTransport::new(transport)); + let mut segments = SessionSegments::new(crate::normalizer::test_session(source)); + receive_expected_request(&mut transport, ClientMessage::Start) + .await? + .ok_or_else(|| anyhow::anyhow!("test transport closed before Start"))?; + + let stream_result = run_started_session(&mut transport, &mut segments).await; + let shutdown_result = segments.into_inner().shutdown().await; + + match (stream_result, shutdown_result) { + (Err(error), _) => Err(error), + (Ok(()), Err(error)) => Err(error), + (Ok(()), Ok(())) => Ok(()), + } +} + +struct DropSignal(Option>); + +impl Drop for DropSignal { + fn drop(&mut self) { + if let Some(sender) = self.0.take() { + let _ = sender.send(()); + } + } +} + +#[tokio::test] +async fn invalid_initial_request_does_not_start_or_poll_source() { + let start_calls = Arc::new(AtomicUsize::new(0)); + let poll_calls = Arc::new(AtomicUsize::new(0)); + let source = { + let start_calls = Arc::clone(&start_calls); + let poll_calls = Arc::clone(&poll_calls); + recording_source(move || { + start_calls.fetch_add(1, Ordering::SeqCst); + async move { + Ok::<_, anyhow::Error>(stream::poll_fn(move |_cx| { + poll_calls.fetch_add(1, Ordering::SeqCst); + Poll::Pending + })) + } + }) + }; + let (transport, client_sender, mut client_receiver) = channel_transport(); + client_sender + .send(Bytes::from_static(b"\x01")) + .expect("send invalid initial Pull"); + let task = tokio::spawn(stream_segments(transport, source, SessionConfig::default())); + + assert_eq!(receive_response(&mut client_receiver).await[0], 2); + assert!(task.await.expect("stream task panicked").is_err()); + assert_eq!(start_calls.load(Ordering::SeqCst), 0); + assert_eq!(poll_calls.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn undecodable_request_while_idle_sends_one_error() { + let source = + recording_source(|| async { Ok::<_, anyhow::Error>(stream::empty::>()) }); + let (transport, client_sender, mut client_receiver) = channel_transport(); + client_sender + .send(Bytes::from_static(b"\x00\x01")) + .expect("send undecodable request"); + let task = tokio::spawn(stream_segments(transport, source, SessionConfig::default())); + + assert_eq!(receive_response(&mut client_receiver).await[0], 2); + assert!(task.await.expect("stream task panicked").is_err()); + assert!(client_receiver.try_recv().is_err()); +} + +#[tokio::test] +async fn transport_error_while_idle_returns_without_response() { + let source = + recording_source(|| async { Ok::<_, anyhow::Error>(stream::empty::>()) }); + let (transport, client_sender, mut client_receiver) = channel_transport(); + client_sender + .send_error(std::io::Error::new( + std::io::ErrorKind::ConnectionReset, + "transport failed", + )) + .expect("send transport error"); + let task = tokio::spawn(stream_segments(transport, source, SessionConfig::default())); + + assert!(task.await.expect("stream task panicked").is_err()); + assert!(client_receiver.try_recv().is_err()); +} + +#[tokio::test] +async fn valid_start_launches_source_before_polling_the_underlying_source() { + let start_calls = Arc::new(AtomicUsize::new(0)); + let (polled_sender, polled_receiver) = oneshot::channel(); + let (release_sender, mut release_receiver) = oneshot::channel(); + let (dropped_sender, dropped_receiver) = oneshot::channel(); + let source = { + let start_calls = Arc::clone(&start_calls); + recording_source(move || { + start_calls.fetch_add(1, Ordering::SeqCst); + async move { + let drop_signal = DropSignal(Some(dropped_sender)); + let mut polled_sender = Some(polled_sender); + let mut emitted = false; + Ok::<_, anyhow::Error>(stream::poll_fn(move |cx| { + let _ = &drop_signal; + if let Some(sender) = polled_sender.take() { + let _ = sender.send(()); + } + if emitted { + return Poll::Ready(None); + } + match Pin::new(&mut release_receiver).poll(cx) { + Poll::Ready(_) => { + emitted = true; + Poll::Ready(Some(Ok(RecordingEvent::SessionEnded))) + } + Poll::Pending => Poll::Pending, + } + })) + } + }) + }; + let (transport, client_sender, mut client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segments(transport, source, SessionConfig::default())); + + assert_eq!(start_calls.load(Ordering::SeqCst), 0); + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + polled_receiver.await.expect("underlying source was polled"); + assert_eq!(start_calls.load(Ordering::SeqCst), 1); + + client_sender.send(Bytes::from_static(b"\x01")).expect("send Pull"); + release_sender.send(()).expect("release source"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x03") + ); + task.await + .expect("stream task panicked") + .expect("stream session failed"); + dropped_receiver.await.expect("running source was dropped"); +} + +#[tokio::test] +async fn aborting_running_session_drops_source() { + let (polled_sender, polled_receiver) = oneshot::channel(); + let (dropped_sender, dropped_receiver) = oneshot::channel(); + let source = recording_source(move || async move { + let drop_signal = DropSignal(Some(dropped_sender)); + let mut polled_sender = Some(polled_sender); + Ok::<_, anyhow::Error>(stream::poll_fn(move |_cx| { + let _ = &drop_signal; + if let Some(sender) = polled_sender.take() { + let _ = sender.send(()); + } + Poll::Pending + })) + }); + let (transport, client_sender, _client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segments(transport, source, SessionConfig::default())); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + polled_receiver.await.expect("underlying source was polled"); + task.abort(); + assert!(task.await.expect_err("aborted stream must not finish").is_cancelled()); + dropped_receiver.await.expect("running source was dropped"); +} + +#[tokio::test] +async fn undecodable_request_waits_for_current_media_and_rejects_only_that_request() { + let (media_waiting_sender, media_waiting_receiver) = oneshot::channel(); + let (media_release_sender, media_release_receiver) = oneshot::channel(); + let source = stream::iter([Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + }))]) + .chain(stream::once(async move { + media_waiting_sender.send(()).expect("signal pending media"); + media_release_receiver.await.expect("release media"); + Ok(SegmentEvent::Data(Bytes::from_static(b"chunk"))) + })) + .chain(stream::pending()); + let (transport, client_sender, mut client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segment_source(transport, source)); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + client_sender.send(Bytes::from_static(b"\x01")).expect("send Pull"); + media_waiting_receiver.await.expect("media was polled"); + client_sender + .send(Bytes::from_static(b"\xFF")) + .expect("send undecodable request"); + client_sender + .send(Bytes::from_static(b"\x01")) + .expect("send unread Pull"); + assert!(client_receiver.try_recv().is_err()); + media_release_sender.send(()).expect("release media"); + + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x00chunk") + ); + assert_eq!(receive_response(&mut client_receiver).await[0], 2); + assert!(task.await.expect("stream task panicked").is_err()); + assert_eq!(client_receiver.recv().await, None); +} + +#[tokio::test] +async fn transport_error_waits_for_current_media_response() { + let (media_waiting_sender, media_waiting_receiver) = oneshot::channel(); + let (media_release_sender, media_release_receiver) = oneshot::channel(); + let source = stream::iter([Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + }))]) + .chain(stream::once(async move { + media_waiting_sender.send(()).expect("signal pending media"); + media_release_receiver.await.expect("release media"); + Ok(SegmentEvent::Data(Bytes::from_static(b"chunk"))) + })) + .chain(stream::pending()); + let (transport, client_sender, mut client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segment_source(transport, source)); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + client_sender.send(Bytes::from_static(b"\x01")).expect("send Pull"); + media_waiting_receiver.await.expect("media was polled"); + client_sender + .send_error(std::io::Error::new( + std::io::ErrorKind::ConnectionReset, + "transport failed", + )) + .expect("send transport error"); + assert!(client_receiver.try_recv().is_err()); + media_release_sender.send(()).expect("release media"); + + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x00chunk") + ); + assert!(task.await.expect("stream task panicked").is_err()); + assert_eq!(client_receiver.recv().await, None); +} + +#[tokio::test] +async fn disconnect_before_start_does_not_start_source() { + let start_calls = Arc::new(AtomicUsize::new(0)); + let source = { + let start_calls = Arc::clone(&start_calls); + recording_source(move || { + start_calls.fetch_add(1, Ordering::SeqCst); + async { Ok::<_, anyhow::Error>(stream::empty::>()) } + }) + }; + let (transport, client_sender, _client_receiver) = channel_transport(); + drop(client_sender); + let task = tokio::spawn(stream_segments(transport, source, SessionConfig::default())); + + task.await + .expect("stream task panicked") + .expect("disconnect should end the stream cleanly"); + assert_eq!(start_calls.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn source_starts_once_after_valid_start() { + let start_calls = Arc::new(AtomicUsize::new(0)); + let source = { + let start_calls = Arc::clone(&start_calls); + recording_source(move || { + start_calls.fetch_add(1, Ordering::SeqCst); + async { Ok::<_, anyhow::Error>(stream::iter([Ok(RecordingEvent::SessionEnded)])) } + }) + }; + let (transport, client_sender, mut client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segments(transport, source, SessionConfig::default())); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + client_sender.send(Bytes::from_static(b"\x01")).expect("send Pull"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x03") + ); + task.await + .expect("stream task panicked") + .expect("stream session failed"); + assert_eq!(start_calls.load(Ordering::SeqCst), 1); +} + +#[tokio::test] +async fn pull_sent_during_launch_is_unread_after_stream_end() { + let (started_sender, started_receiver) = oneshot::channel(); + let (release_sender, release_receiver) = oneshot::channel(); + let source = recording_source(move || async move { + started_sender.send(()).expect("signal startup"); + release_receiver.await.expect("release startup"); + Ok::<_, anyhow::Error>(stream::iter([Ok(RecordingEvent::SessionEnded)])) + }); + let (transport, client_sender, mut client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segments(transport, source, SessionConfig::default())); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + client_sender + .send(Bytes::from_static(b"\x01")) + .expect("send unread Pull"); + started_receiver.await.expect("startup was polled"); + assert!( + client_receiver.try_recv().is_err(), + "launch must not respond before the source is ready" + ); + release_sender.send(()).expect("release startup"); + + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x03") + ); + task.await + .expect("stream task panicked") + .expect("stream session failed"); + assert_eq!(client_receiver.recv().await, None); +} + +#[tokio::test] +async fn startup_failure_rejects_the_accepted_start_only() { + let (started_sender, started_receiver) = oneshot::channel(); + let (release_sender, release_receiver) = oneshot::channel(); + let source = recording_source(move || async move { + started_sender.send(()).expect("signal startup"); + release_receiver.await.expect("release startup"); + Err::>, _>(anyhow::anyhow!("startup failed")) + }); + let (transport, client_sender, mut client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segments(transport, source, SessionConfig::default())); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + client_sender + .send(Bytes::from_static(b"\x01")) + .expect("send unread Pull"); + started_receiver.await.expect("startup was polled"); + release_sender.send(()).expect("release startup"); + + assert_eq!(receive_response(&mut client_receiver).await[0], 2); + assert!( + client_receiver.try_recv().is_err(), + "unread Pull must not receive a startup error" + ); + assert!(task.await.expect("stream task panicked").is_err()); + assert_eq!(client_receiver.recv().await, None); +} + +#[tokio::test] +async fn disconnect_waits_for_current_media_response() { + let (media_waiting_sender, media_waiting_receiver) = oneshot::channel(); + let (media_release_sender, media_release_receiver) = oneshot::channel(); + let source = stream::iter([Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + }))]) + .chain(stream::once(async move { + media_waiting_sender.send(()).expect("signal pending media"); + media_release_receiver.await.expect("release media"); + Ok(SegmentEvent::Data(Bytes::from_static(b"chunk"))) + })) + .chain(stream::pending()); + let (transport, client_sender, mut client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segment_source(transport, source)); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + client_sender.send(Bytes::from_static(b"\x01")).expect("send Pull"); + media_waiting_receiver.await.expect("media was polled"); + drop(client_sender); + assert!(client_receiver.try_recv().is_err()); + media_release_sender.send(()).expect("release media"); + + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x00chunk") + ); + task.await + .expect("stream task panicked") + .expect("disconnect should end the stream cleanly"); + assert_eq!(client_receiver.recv().await, None); +} + +#[tokio::test] +async fn abort_during_startup_drops_pending_source() { + let (dropped_sender, dropped_receiver) = oneshot::channel(); + let (started_sender, started_receiver) = oneshot::channel(); + let source = recording_source(move || async move { + let _drop_signal = DropSignal(Some(dropped_sender)); + started_sender.send(()).expect("signal startup"); + futures_util::future::pending::<()>().await; + Ok::<_, anyhow::Error>(stream::empty::>()) + }); + let (transport, client_sender, _client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segments(transport, source, SessionConfig::default())); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + started_receiver.await.expect("startup was polled"); + task.abort(); + assert!(task.await.expect_err("aborted stream must not finish").is_cancelled()); + dropped_receiver.await.expect("pending source was dropped"); +} + +#[test] +fn protocol_codes_are_stable() { + assert_eq!( + encode_server_message(ServerMessage::Metadata), + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + assert_eq!( + encode_server_message(ServerMessage::SegmentStarted), + Bytes::from_static(b"\x04{\"codec\":\"vp8\"}") + ); + assert_eq!( + encode_server_message(ServerMessage::Chunk(Bytes::from_static(b"webm"))), + Bytes::from_static(b"\x00webm") + ); + assert_eq!( + encode_server_message(ServerMessage::StreamEnded), + Bytes::from_static(b"\x03") + ); +} + +#[test] +fn client_messages_require_one_complete_transport_message() { + assert_eq!( + decode_client_message(b"\x00").expect("decode start"), + ClientMessage::Start + ); + assert_eq!( + decode_client_message(b"\x01").expect("decode pull"), + ClientMessage::Pull + ); + assert!(decode_client_message(b"\x02").is_err()); + assert!(decode_client_message(b"\x00\x01").is_err()); + assert!(decode_client_message(b"").is_err()); +} + +#[tokio::test] +async fn transport_adapts_typed_messages_at_the_wire_boundary() { + let (transport, client_sender, mut client_receiver) = channel_transport(); + let mut transport = SessionTransport::new(CodecTransport::new(transport)); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + assert_eq!( + transport + .recv() + .await + .expect("transport message") + .expect("decoded client message"), + ClientMessage::Start + ); + + transport + .send(ServerMessage::StreamEnded) + .await + .expect("send typed server message"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x03") + ); +} + +#[tokio::test] +async fn segment_end_is_implicit_on_the_wire() { + let events = [ + Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + })), + Ok(SegmentEvent::Data(Bytes::from_static(b"first"))), + Ok(SegmentEvent::End), + Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 1, + width: 800, + height: 600, + })), + Ok(SegmentEvent::Data(Bytes::from_static(b"second"))), + Ok(SegmentEvent::End), + ]; + let mut segments = SessionSegments::new(stream::iter(events)); + + assert_eq!( + segments.next().await.expect("first data"), + ServerMessage::Chunk(Bytes::from_static(b"first")) + ); + assert_eq!( + segments.next().await.expect("second begin"), + ServerMessage::SegmentStarted + ); + assert_eq!( + segments.next().await.expect("second data"), + ServerMessage::Chunk(Bytes::from_static(b"second")) + ); + assert_eq!(segments.next().await.expect("stream end"), ServerMessage::StreamEnded); +} + +#[tokio::test] +async fn multi_segment_protocol_transcript_is_stable() { + let source = segment_source([ + Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + })), + Ok(SegmentEvent::Data(Bytes::from_static(b"first"))), + Ok(SegmentEvent::End), + Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 1, + width: 800, + height: 600, + })), + Ok(SegmentEvent::Data(Bytes::from_static(b"second"))), + Ok(SegmentEvent::End), + ]); + let (transport, client_sender, mut client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segment_source(transport, stream::iter(source))); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + + for expected in [ + Bytes::from_static(b"\x00first"), + Bytes::from_static(b"\x04{\"codec\":\"vp8\"}"), + Bytes::from_static(b"\x00second"), + Bytes::from_static(b"\x03"), + ] { + client_sender.send(Bytes::from_static(b"\x01")).expect("send Pull"); + assert_eq!(receive_response(&mut client_receiver).await, expected); + } + + task.await + .expect("stream task panicked") + .expect("stream session failed"); +} + +#[tokio::test] +async fn buffered_pulls_are_served_in_order() { + let (transport, client_sender, mut client_receiver) = channel_transport(); + let source = segment_source([ + Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + })), + Ok(SegmentEvent::Data(Bytes::from_static(b"first"))), + Ok(SegmentEvent::Data(Bytes::from_static(b"second"))), + Ok(SegmentEvent::End), + ]); + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + for _ in 0..3 { + client_sender.send(Bytes::from_static(b"\x01")).expect("send Pull"); + } + let task = tokio::spawn(stream_segment_source(transport, stream::iter(source))); + + assert_eq!(receive_response(&mut client_receiver).await[0], 1); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x00first") + ); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x00second") + ); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x03") + ); + task.await + .expect("stream task panicked") + .expect("stream session failed"); +} + +#[tokio::test] +async fn buffered_pull_is_read_after_the_current_media_response() { + let (media_waiting_sender, media_waiting_receiver) = oneshot::channel(); + let (media_release_sender, media_release_receiver) = oneshot::channel(); + let source = stream::iter([Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + }))]) + .chain(stream::once(async move { + media_waiting_sender.send(()).expect("signal pending media"); + media_release_receiver.await.expect("release media"); + Ok(SegmentEvent::Data(Bytes::from_static(b"chunk"))) + })) + .chain(stream::iter([Ok(SegmentEvent::End)])); + let (transport, client_sender, mut client_receiver) = channel_transport(); + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + let task = tokio::spawn(stream_segment_source(transport, source)); + + assert_eq!(receive_response(&mut client_receiver).await[0], 1); + client_sender.send(Bytes::from_static(b"\x01")).expect("send Pull"); + media_waiting_receiver.await.expect("media was polled"); + client_sender + .send(Bytes::from_static(b"\x01")) + .expect("send buffered Pull"); + assert!(client_receiver.try_recv().is_err()); + media_release_sender.send(()).expect("release media"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x00chunk") + ); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x03") + ); + task.await + .expect("stream task panicked") + .expect("stream session failed"); +} + +#[tokio::test] +async fn wrong_state_request_waits_for_current_media_and_sends_one_error() { + let (media_waiting_sender, media_waiting_receiver) = oneshot::channel(); + let (media_release_sender, media_release_receiver) = oneshot::channel(); + let source = stream::iter([Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + }))]) + .chain(stream::once(async move { + media_waiting_sender.send(()).expect("signal pending media"); + media_release_receiver.await.expect("release media"); + Ok(SegmentEvent::Data(Bytes::from_static(b"chunk"))) + })) + .chain(stream::pending()); + let (transport, client_sender, mut client_receiver) = channel_transport(); + let task = tokio::spawn(stream_segment_source(transport, source)); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + client_sender.send(Bytes::from_static(b"\x01")).expect("send Pull"); + media_waiting_receiver.await.expect("media was polled"); + client_sender + .send(Bytes::from_static(b"\x00")) + .expect("send Start in running state"); + client_sender + .send(Bytes::from_static(b"\x01")) + .expect("send unread Pull"); + assert!(client_receiver.try_recv().is_err()); + media_release_sender.send(()).expect("release media"); + + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x00chunk") + ); + assert_eq!(receive_response(&mut client_receiver).await[0], 2); + assert!(task.await.expect("stream task panicked").is_err()); + assert_eq!(client_receiver.recv().await, None); +} + +#[tokio::test] +async fn first_segment_sequence_must_be_zero() { + let mut segments = SessionSegments::new(stream::iter([Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 1, + width: 640, + height: 480, + }))])); + + let error = segments.next().await.expect_err("nonzero first sequence must fail"); + + assert!( + format!("{error:#}").contains("segment sequence is not contiguous"), + "{error:#}" + ); +} + +#[tokio::test] +async fn segment_sequence_gap_is_rejected() { + let mut segments = SessionSegments::new(stream::iter([ + Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + })), + Ok(SegmentEvent::End), + Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 2, + width: 800, + height: 600, + })), + ])); + + let error = segments.next().await.expect_err("segment sequence gap must fail"); + + assert!( + format!("{error:#}").contains("segment sequence is not contiguous"), + "{error:#}" + ); +} + +#[tokio::test] +async fn stream_end_answers_current_request_only() { + let (transport, client_sender, mut client_receiver) = channel_transport(); + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + client_sender + .send(Bytes::from_static(b"\x01")) + .expect("send unread Pull"); + let source = stream::empty::>(); + let task = tokio::spawn(stream_segment_source(transport, source)); + + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x03") + ); + task.await + .expect("stream task panicked") + .expect("stream session failed"); + assert_eq!(client_receiver.recv().await, None); +} + +#[tokio::test] +async fn each_request_receives_exactly_one_response() { + let (transport, client_sender, mut client_receiver) = channel_transport(); + let source = segment_source([ + Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + })), + Ok(SegmentEvent::Data(Bytes::from_static(b"chunk"))), + Ok(SegmentEvent::End), + ]); + let task = tokio::spawn(stream_segment_source(transport, stream::iter(source))); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + assert_eq!(receive_response(&mut client_receiver).await[0], 1); + assert!( + client_receiver.try_recv().is_err(), + "server sent a response without another request" + ); + + client_sender + .send(Bytes::from_static(b"\x01")) + .expect("send first Pull"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x00chunk") + ); + assert!( + client_receiver.try_recv().is_err(), + "server sent a response without another request" + ); + + client_sender + .send(Bytes::from_static(b"\x01")) + .expect("send final Pull"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x03") + ); + task.await + .expect("stream task panicked") + .expect("stream session failed"); +} + +#[tokio::test] +async fn segment_failure_sends_one_error_for_the_current_request() { + let (transport, client_sender, mut client_receiver) = channel_transport(); + let source = segment_source([Err(anyhow::anyhow!("test segment failure"))]); + let task = tokio::spawn(stream_segment_source(transport, stream::iter(source))); + + client_sender.send(Bytes::from_static(b"\x00")).expect("send Start"); + client_sender + .send(Bytes::from_static(b"\x01")) + .expect("send unread Pull"); + assert_eq!( + receive_response(&mut client_receiver).await, + Bytes::from_static(b"\x01{\"codec\":\"vp8\"}") + ); + let response = receive_response(&mut client_receiver).await; + assert_eq!(response[0], 2); + assert!(task.await.expect("stream task panicked").is_err()); + assert_eq!( + client_receiver.recv().await, + None, + "error must not be followed by StreamEnded" + ); +} diff --git a/crates/video-streamer/src/protocol/transport.rs b/crates/video-streamer/src/protocol/transport.rs new file mode 100644 index 000000000..e252434e0 --- /dev/null +++ b/crates/video-streamer/src/protocol/transport.rs @@ -0,0 +1,106 @@ +use std::error::Error; +use std::pin::Pin; +use std::task::{Context, Poll}; + +use anyhow::Context as _; +use bytes::Bytes; +use futures_util::{Sink, SinkExt as _, Stream, StreamExt as _}; + +use super::message::{ClientMessage, ServerMessage, UserFriendlyError, decode_client_message, encode_server_message}; + +#[derive(Debug)] +pub(super) enum ReceiveError { + Transport(E), + Decode(anyhow::Error), +} + +pub(super) struct CodecTransport { + inner: T, +} + +impl CodecTransport { + pub(super) fn new(inner: T) -> Self { + Self { inner } + } +} + +impl Stream for CodecTransport +where + T: Stream> + Unpin, + E: Error + Send + Sync + 'static, +{ + type Item = Result>; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_next(cx).map(|incoming| { + incoming.map(|result| match result { + Ok(bytes) => decode_client_message(&bytes).map_err(ReceiveError::Decode), + Err(error) => Err(ReceiveError::Transport(error)), + }) + }) + } +} + +impl Sink for CodecTransport +where + T: Sink + Unpin, + E: Error + Send + Sync + 'static, +{ + type Error = E; + + fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_ready(cx) + } + + fn start_send(self: Pin<&mut Self>, message: ServerMessage) -> Result<(), Self::Error> { + Pin::new(&mut self.get_mut().inner).start_send(encode_server_message(message)) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_flush(cx) + } + + fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_close(cx) + } +} + +pub(super) struct SessionTransport { + inner: T, +} + +impl SessionTransport { + pub(super) fn new(inner: T) -> Self { + Self { inner } + } +} + +impl SessionTransport +where + T: Stream>> + Unpin, + E: Error + Send + Sync + 'static, +{ + pub(super) async fn recv(&mut self) -> Option>> { + self.inner.next().await + } +} + +impl SessionTransport +where + T: Sink + Unpin, + E: Error + Send + Sync + 'static, +{ + pub(super) async fn send(&mut self, message: ServerMessage) -> anyhow::Result<()> { + self.inner + .send(message) + .await + .map_err(anyhow::Error::new) + .context("write server stream message") + } + + pub(super) async fn reject(&mut self) { + let _ = self + .send(ServerMessage::Error(UserFriendlyError::UnexpectedError)) + .await; + } +} diff --git a/crates/video-streamer/src/session.rs b/crates/video-streamer/src/session.rs new file mode 100644 index 000000000..111e54abd --- /dev/null +++ b/crates/video-streamer/src/session.rs @@ -0,0 +1,99 @@ +use std::error::Error; +use std::fmt; +use std::future::Future; +use std::io::{self, Read, Seek, SeekFrom}; + +use bytes::Bytes; +use futures_util::{Sink, Stream}; + +trait RecordingClipReader: Read + Seek + Send {} + +impl RecordingClipReader for T where T: Read + Seek + Send {} + +/// Owns the seekable reader for one recording clip. +/// +/// The reader remains attached to this clip while consumers seek for bounded replay. +/// Replay uses this owned handle instead of reopening the clip path. +pub struct RecordingClip { + reader: Box, +} + +impl RecordingClip { + pub fn new(reader: R) -> Self + where + R: Read + Seek + Send + 'static, + { + Self { + reader: Box::new(reader), + } + } + + pub(crate) fn read(&mut self, buffer: &mut [u8]) -> io::Result { + self.reader.read(buffer) + } + + pub(crate) fn seek(&mut self, position: SeekFrom) -> io::Result { + self.reader.seek(position) + } +} + +impl fmt::Debug for RecordingClip { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.debug_struct("RecordingClip").finish_non_exhaustive() + } +} + +/// A structural event from one append-only recording session. +#[derive(Debug)] +pub enum RecordingEvent { + ClipStarted { + sequence: u64, + start_at: StartAt, + clip: RecordingClip, + }, + DataAvailable, + CaughtUp, + ClipEnded, + SessionEnded, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum StartAt { + Beginning, + LiveEdge, +} + +#[derive(Clone, Copy, Debug)] +pub struct SessionConfig { + pub encoder_threads: u32, + pub adaptive_frame_skip: bool, +} + +impl Default for SessionConfig { + fn default() -> Self { + Self { + encoder_threads: u32::try_from(num_cpus::get()).unwrap_or(1).max(1), + adaptive_frame_skip: true, + } + } +} + +/// Produces structural events and transfers each clip reader to the consumer once. +pub trait RecordingSource: Send + 'static { + type Stream: Stream> + Send + 'static; + type Start: Future> + Send + 'static; + + fn start(self) -> Self::Start; +} + +/// Converts a recording session into independent VP8 WebM segments over one pull-driven stream. +/// +/// Each segment has one resolution, and output sequence numbers remain contiguous across input clips. +pub async fn stream_session(source: S, transport: T, config: SessionConfig) -> anyhow::Result<()> +where + S: RecordingSource, + T: Stream> + Sink + Unpin, + E: Error + Send + Sync + 'static, +{ + crate::protocol::stream_segments(transport, source, config).await +} diff --git a/crates/video-streamer/src/streamer/block_tag.rs b/crates/video-streamer/src/streamer/block_tag.rs index 7ee0a9f23..98d1238b5 100644 --- a/crates/video-streamer/src/streamer/block_tag.rs +++ b/crates/video-streamer/src/streamer/block_tag.rs @@ -12,6 +12,7 @@ pub(crate) enum BlockTag { #[derive(Clone)] pub(crate) struct VideoBlock { + pub(crate) track: u64, pub(crate) cluster_timestamp: Option, pub(crate) timestamp: i16, pub(crate) is_key_frame: bool, @@ -22,6 +23,7 @@ impl fmt::Debug for VideoBlock { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { f.debug_struct("VideoBlock") .field("cluster_timestamp", &self.cluster_timestamp) + .field("track", &self.track) .field("timestamp", &self.timestamp) .field("is_key_frame", &self.is_key_frame) .field( @@ -58,6 +60,7 @@ impl VideoBlock { .any(|frame| is_vpx_key_frame(frame.data, codec)); Self { + track: block.track, cluster_timestamp, block_tag: BlockTag::BlockGroup(children), timestamp, @@ -67,6 +70,7 @@ impl VideoBlock { MatroskaSpec::SimpleBlock(data) => { let simple_block = SimpleBlock::try_from(&data)?; Self { + track: simple_block.track, cluster_timestamp, timestamp: simple_block.timestamp, is_key_frame: simple_block.keyframe, @@ -80,14 +84,15 @@ impl VideoBlock { } pub(crate) fn absolute_timestamp(&self) -> anyhow::Result { - let timestamp = u64::try_from(self.timestamp)?; - Ok(self + let cluster_timestamp = self .cluster_timestamp - .with_context(|| format!("Cluster timestamp not found for timestamp: {}", self.timestamp))? - + timestamp) + .with_context(|| format!("Cluster timestamp not found for timestamp: {}", self.timestamp))?; + let timestamp = i64::try_from(cluster_timestamp)? + .checked_add(i64::from(self.timestamp)) + .context("block timestamp overflow")?; + u64::try_from(timestamp).context("negative absolute block timestamp") } - // We only handle non-lacing frames for now pub(crate) fn get_frame(&self) -> anyhow::Result> { let mut frames: Vec<_> = match &self.block_tag { BlockTag::SimpleBlock(data) => { @@ -120,8 +125,8 @@ impl VideoBlock { } }; - assert!(frames.len() == 1); - Ok(frames.pop().expect("frame length was asserted")) + anyhow::ensure!(frames.len() == 1, "laced video blocks are not supported"); + frames.pop().context("video block contains no frame") } } @@ -225,6 +230,20 @@ mod tests { ); } + #[test] + fn get_frame_rejects_laced_blocks() { + let tag = MatroskaSpec::SimpleBlock(vec![0x81, 0x00, 0x01, 0x84, 0x01, 0x00, 0x00]); + let video_block = VideoBlock::new(tag, None, VpxCodec::VP8).expect("laced block should parse"); + + assert_eq!( + video_block + .get_frame() + .expect_err("laced block should be rejected") + .to_string(), + "laced video blocks are not supported" + ); + } + #[test] fn vp9_empty_buffer_is_not_keyframe() { assert!(!is_vp9_key_frame(&[])); diff --git a/crates/video-streamer/src/streamer/signal_writer.rs b/crates/video-streamer/src/streamer/signal_writer.rs index e66af86ff..21bdac1cd 100644 --- a/crates/video-streamer/src/streamer/signal_writer.rs +++ b/crates/video-streamer/src/streamer/signal_writer.rs @@ -23,7 +23,11 @@ where cx: &mut std::task::Context<'_>, buf: &[u8], ) -> Poll> { - tokio::io::AsyncWrite::poll_write(std::pin::Pin::new(&mut self.writer), cx, buf) + let result = tokio::io::AsyncWrite::poll_write(std::pin::Pin::new(&mut self.writer), cx, buf); + if matches!(&result, Poll::Ready(Ok(written)) if *written > 0) { + self.notify.notify_one(); + } + result } fn poll_flush( @@ -34,7 +38,7 @@ where return Poll::Pending; }; - self.notify.notify_waiters(); + self.notify.notify_one(); Poll::Ready(res) } diff --git a/devolutions-gateway/src/api/jrec.rs b/devolutions-gateway/src/api/jrec.rs index 5212eb095..83c1709d3 100644 --- a/devolutions-gateway/src/api/jrec.rs +++ b/devolutions-gateway/src/api/jrec.rs @@ -965,7 +965,11 @@ impl From for CloseFrame { } async fn shadow_recording( - State(DgwState { recordings, .. }): State, + State(DgwState { + recordings, + shutdown_signal, + .. + }): State, extract::Path(id): extract::Path, JrecToken(claims): JrecToken, ws: WebSocketUpgrade, @@ -978,31 +982,22 @@ async fn shadow_recording( return close_with_error(ws, StreamerCloseCode::StreamingEnded); } - let Ok(Some(crate::recording::OnGoingRecordingState::Connected)) = recordings.get_state(id).await else { - return close_with_error(ws, StreamerCloseCode::StreamingEnded); - }; - if !xmf::is_init() { warn!(%id, "Shadow recording rejected: XMF native library is not loaded"); return close_with_error(ws, StreamerCloseCode::InternalError); } - let Ok(notify) = recordings.subscribe_to_recording_finish(id).await else { - warn!(%id, "Shadow recording rejected: failed to subscribe to recording finish"); - return close_with_error(ws, StreamerCloseCode::InternalError); - }; - let Ok(recording_files) = recordings.list_files(id).await else { warn!(%id, "Shadow recording rejected: failed to list recording files"); return close_with_error(ws, StreamerCloseCode::InternalError); }; - let Some(recording_path) = recording_files.last() else { + if recording_files.is_empty() { warn!(%id, "Shadow recording rejected: no recording files found"); return close_with_error(ws, StreamerCloseCode::InternalError); - }; + } - return crate::streaming::stream_file(recording_path, ws, notify, recordings, id) + return crate::streaming::stream_recording(ws, shutdown_signal, recordings, id) .await .map_err(|_| HttpError::internal().msg("failed to stream file")); diff --git a/devolutions-gateway/src/recording.rs b/devolutions-gateway/src/recording.rs index f07913868..6283840b9 100644 --- a/devolutions-gateway/src/recording.rs +++ b/devolutions-gateway/src/recording.rs @@ -14,7 +14,7 @@ use futures::future::Either; use parking_lot::Mutex; use serde::Serialize; use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt, BufWriter}; -use tokio::sync::{Notify, mpsc, oneshot}; +use tokio::sync::{mpsc, oneshot, watch}; use tokio::{fs, io}; use typed_builder::TypedBuilder; use uuid::Uuid; @@ -132,6 +132,7 @@ where let res = match open_options.open(&recording_file).await { Ok(file) => { + recordings.clip_started(session_id).await?; // Wrap SignalWriter inside a BufWriter to reduce the number of flushes. let (file, flush_signal) = SignalWriter::new(file); // larger buffer size to reduce the number of flushes @@ -144,7 +145,7 @@ where loop { tokio::select! { _ = flush_signal.notified() => { - recordings.new_chunk_appended(session_id)?; + recordings.new_chunk_appended(session_id).await?; }, _ = shutdown_signal_clone.wait() => { break; @@ -173,8 +174,22 @@ where }; signal_loop.abort(); + let _ = signal_loop.await; - res + let flush_result = file.flush().await; + if flush_result.is_ok() { + recordings.new_chunk_appended(session_id).await?; + } + + match (res, flush_result) { + (Err(error), _) => Err(error), + (Ok(_), Err(error)) if is_storage_full(&error) => { + warn!(%session_id, "Recording storage is full; closing push stream"); + Ok(PushOutcome::StorageFull) + } + (Ok(_), Err(error)) => Err(anyhow::Error::new(error).context("flush JREC recording file")), + (Ok(outcome), Ok(())) => Ok(outcome), + } } Err(e) => Err(anyhow::Error::new(e).context(format!("failed to open file at {recording_file}"))), }; @@ -241,6 +256,76 @@ struct OnGoingRecording { manifest_path: Utf8PathBuf, session_must_be_recorded: bool, disconnected_ttl: Duration, + stream_state: watch::Sender, +} + +#[derive(Clone, Debug)] +pub(crate) struct RecordingStreamClip { + pub(crate) sequence: u64, + pub(crate) path: Utf8PathBuf, +} + +#[derive(Clone, Copy, Debug)] +pub(crate) struct ActiveRecordingStreamClip { + pub(crate) sequence: u64, + pub(crate) ready: bool, +} + +#[derive(Clone, Debug)] +pub(crate) struct RecordingStreamState { + pub(crate) clips: Arc>, + pub(crate) active: Option, + pub(crate) ended: bool, +} + +impl RecordingStreamState { + pub(crate) fn mark_disconnected(&mut self) { + self.active = None; + self.ended = false; + } + + pub(crate) fn mark_ended(&mut self) { + self.active = None; + self.ended = true; + } + + #[cfg(test)] + pub(crate) fn for_test( + clips: Vec, + active: Option, + ended: bool, + ) -> Self { + Self { + clips: Arc::new(clips), + active, + ended, + } + } +} + +#[cfg(test)] +mod stream_state_tests { + use super::*; + + #[test] + fn disconnect_is_not_a_confirmed_session_end() { + let mut state = RecordingStreamState { + clips: Arc::new(Vec::new()), + active: Some(ActiveRecordingStreamClip { + sequence: 0, + ready: true, + }), + ended: false, + }; + + state.mark_disconnected(); + assert!(state.active.is_none()); + assert!(!state.ended); + + state.mark_ended(); + assert!(state.active.is_none()); + assert!(state.ended); + } } enum RecordingManagerMessage { @@ -253,6 +338,12 @@ enum RecordingManagerMessage { Disconnect { id: Uuid, }, + ClipStarted { + id: Uuid, + }, + ChunkAppended { + id: Uuid, + }, GetState { id: Uuid, channel: oneshot::Sender>, @@ -268,9 +359,9 @@ enum RecordingManagerMessage { id: Uuid, session_must_be_recorded: bool, }, - SubscribeToSessionEndNotification { + SubscribeToStream { id: Uuid, - channel: oneshot::Sender>, + channel: oneshot::Sender>, }, } @@ -289,6 +380,8 @@ impl fmt::Debug for RecordingManagerMessage { .field("disconnected_ttl", disconnected_ttl) .finish_non_exhaustive(), RecordingManagerMessage::Disconnect { id } => f.debug_struct("Disconnect").field("id", id).finish(), + RecordingManagerMessage::ClipStarted { id } => f.debug_struct("ClipStarted").field("id", id).finish(), + RecordingManagerMessage::ChunkAppended { id } => f.debug_struct("ChunkAppended").field("id", id).finish(), RecordingManagerMessage::GetState { id, channel: _ } => { f.debug_struct("GetState").field("id", id).finish_non_exhaustive() } @@ -301,12 +394,13 @@ impl fmt::Debug for RecordingManagerMessage { .field("id", id) .field("session_must_be_recorded", session_must_be_recorded) .finish(), - RecordingManagerMessage::SubscribeToSessionEndNotification { id, channel: _ } => { - f.debug_struct("SubscribeToOngoingRecording").field("id", id).finish() - } RecordingManagerMessage::ListFiles { id, channel: _ } => { f.debug_struct("ListFiles").field("id", id).finish() } + RecordingManagerMessage::SubscribeToStream { id, channel: _ } => f + .debug_struct("SubscribeToStream") + .field("id", id) + .finish_non_exhaustive(), } } } @@ -386,24 +480,37 @@ impl RecordingMessageSender { senders.push(tx); } - pub(crate) fn new_chunk_appended(&self, recording_id: Uuid) -> anyhow::Result<()> { - let senders = { self.flush_map.lock().remove(&recording_id) }; + async fn clip_started(&self, recording_id: Uuid) -> anyhow::Result<()> { + self.channel + .send(RecordingManagerMessage::ClipStarted { id: recording_id }) + .await + .ok() + .context("couldn't send ClipStarted message") + } - let Some(senders) = senders else { - return Ok(()); - }; + pub(crate) async fn new_chunk_appended(&self, recording_id: Uuid) -> anyhow::Result<()> { + let senders = { self.flush_map.lock().remove(&recording_id) }; - for tx in senders { - let _ = tx.send(()); + if let Some(senders) = senders { + for tx in senders { + let _ = tx.send(()); + } } - Ok(()) + self.channel + .send(RecordingManagerMessage::ChunkAppended { id: recording_id }) + .await + .ok() + .context("couldn't send ChunkAppended message") } - pub(crate) async fn subscribe_to_recording_finish(&self, recording_id: Uuid) -> anyhow::Result> { + pub(crate) async fn subscribe_to_stream( + &self, + recording_id: Uuid, + ) -> anyhow::Result> { let (tx, rx) = oneshot::channel(); self.channel - .send(RecordingManagerMessage::SubscribeToSessionEndNotification { + .send(RecordingManagerMessage::SubscribeToStream { id: recording_id, channel: tx, }) @@ -479,7 +586,6 @@ impl Ord for DisconnectedTtl { pub struct RecordingManagerTask { rx: RecordingMessageReceiver, ongoing_recordings: HashMap, - recording_end_notifier: HashMap>, recordings_path: Utf8PathBuf, session_manager_handle: SessionMessageSender, job_queue_handle: JobQueueHandle, @@ -495,7 +601,6 @@ impl RecordingManagerTask { Self { rx, ongoing_recordings: HashMap::new(), - recording_end_notifier: HashMap::new(), recordings_path, session_manager_handle, job_queue_handle, @@ -516,6 +621,10 @@ impl RecordingManagerTask { anyhow::bail!("concurrent recording for the same session is not supported"); } + let existing_stream_state = self + .ongoing_recordings + .get(&id) + .map(|ongoing| ongoing.stream_state.clone()); let recording_path = self.recordings_path.join(id.to_string()); let manifest_path = recording_path.join("recording.json"); @@ -588,6 +697,43 @@ impl RecordingManagerTask { .map(|info| info.recording_policy) .unwrap_or(false); + let sequence = manifest + .files + .len() + .checked_sub(1) + .context("recording manifest has no files")?; + let sequence = u64::try_from(sequence).context("recording sequence does not fit in u64")?; + let clip = RecordingStreamClip { + sequence, + path: recording_file.clone(), + }; + let stream_state = if let Some(stream_state) = existing_stream_state { + stream_state.send_modify(|state| { + Arc::make_mut(&mut state.clips).push(clip.clone()); + state.active = Some(ActiveRecordingStreamClip { sequence, ready: false }); + state.ended = false; + }); + stream_state + } else { + let clips = manifest + .files + .iter() + .enumerate() + .map(|(sequence, file)| { + Ok(RecordingStreamClip { + sequence: u64::try_from(sequence).context("recording sequence does not fit in u64")?, + path: recording_path.join(&file.file_name), + }) + }) + .collect::>>()?; + let state = RecordingStreamState { + clips: Arc::new(clips), + active: Some(ActiveRecordingStreamClip { sequence, ready: false }), + ended: false, + }; + watch::channel(state).0 + }; + self.ongoing_recordings.insert( id, OnGoingRecording { @@ -596,6 +742,7 @@ impl RecordingManagerTask { manifest_path, session_must_be_recorded, disconnected_ttl, + stream_state, }, ); let ongoing_recording_count = self.ongoing_recordings.len(); @@ -612,6 +759,51 @@ impl RecordingManagerTask { Ok(recording_file) } + fn handle_clip_started(&mut self, id: Uuid) -> anyhow::Result<()> { + let ongoing = self + .ongoing_recordings + .get(&id) + .with_context(|| format!("unknown recording for ID {id}"))?; + let active = ongoing + .stream_state + .borrow() + .active + .context("recording has no active clip")?; + + if !matches!(ongoing.state, OnGoingRecordingState::Connected) || active.ready { + anyhow::bail!("recording clip can’t be started in its current state"); + } + + ongoing.stream_state.send_modify(|state| { + state.active = Some(ActiveRecordingStreamClip { + sequence: active.sequence, + ready: true, + }); + }); + + Ok(()) + } + + fn handle_chunk_appended(&mut self, id: Uuid) -> anyhow::Result<()> { + let ongoing = self + .ongoing_recordings + .get(&id) + .with_context(|| format!("unknown recording for ID {id}"))?; + let active = ongoing + .stream_state + .borrow() + .active + .context("recording has no active clip")?; + + if !active.ready { + anyhow::bail!("recording clip is not ready"); + } + + ongoing.stream_state.send_modify(|_| {}); + + Ok(()) + } + async fn handle_disconnect(&mut self, id: Uuid) -> anyhow::Result<()> { let Some(ongoing) = self.ongoing_recordings.get_mut(&id) else { return Err(anyhow::anyhow!("unknown recording for ID {id}")); @@ -647,10 +839,9 @@ impl RecordingManagerTask { .save_to_file(&ongoing.manifest_path) .with_context(|| format!("write manifest at {}", ongoing.manifest_path))?; - // Notify all the streamers that recording has ended. - if let Some(notify) = self.recording_end_notifier.get(&id) { - notify.notify_waiters(); - } + ongoing + .stream_state + .send_modify(RecordingStreamState::mark_disconnected); info!(%id, "Start video remuxing operation"); if recording_file_path.extension() == Some(RecordingFileType::WebM.extension()) { @@ -686,6 +877,7 @@ impl RecordingManagerTask { OnGoingRecordingState::LastSeen { timestamp } if now >= timestamp + disconnected_ttl_secs - 1 => { debug!(%id, "Mark recording as terminated"); self.rx.active_recordings.remove(id); + ongoing.stream_state.send_modify(RecordingStreamState::mark_ended); // Check the recording policy of the associated session and kill it if necessary. if ongoing.session_must_be_recorded { @@ -722,7 +914,6 @@ impl RecordingManagerTask { } self.ongoing_recordings.remove(&id); - self.recording_end_notifier.remove(&id); } _ => { trace!(%id, "Recording should not be removed yet"); @@ -731,19 +922,12 @@ impl RecordingManagerTask { } } - fn subscribe(&mut self, id: Uuid) -> anyhow::Result> { - debug!(%id, "Subscribing to ongoing recording"); - if !self.ongoing_recordings.contains_key(&id) { - anyhow::bail!("unknown recording for ID {id}"); - } - - if let Some(notify) = self.recording_end_notifier.get(&id) { - Ok(Arc::clone(notify)) - } else { - let notify = Arc::new(Notify::new()); - self.recording_end_notifier.insert(id, Arc::clone(¬ify)); - Ok(notify) - } + fn subscribe_stream(&self, id: Uuid) -> anyhow::Result> { + let ongoing = self + .ongoing_recordings + .get(&id) + .with_context(|| format!("unknown recording for ID {id}"))?; + Ok(ongoing.stream_state.subscribe()) } } @@ -822,6 +1006,16 @@ async fn recording_manager_task( } } } + RecordingManagerMessage::ClipStarted { id } => { + if let Err(error) = manager.handle_clip_started(id) { + error!(%error, "handle_clip_started"); + } + } + RecordingManagerMessage::ChunkAppended { id } => { + if let Err(error) = manager.handle_chunk_appended(id) { + error!(%error, "handle_chunk_appended"); + } + } RecordingManagerMessage::GetState { id, channel } => { let response = manager.ongoing_recordings.get(&id).map(|ongoing| ongoing.state.clone()); let _ = channel.send(response); @@ -839,14 +1033,14 @@ async fn recording_manager_task( ); } }, - RecordingManagerMessage::SubscribeToSessionEndNotification {id, channel } => { - match manager.subscribe(id) { - Ok(notifier) => { - let _ = channel.send(notifier); - }, - Err(e) => error!(error = format!("{e:#}"), "subscribe to session end notification"), + RecordingManagerMessage::SubscribeToStream { id, channel } => { + match manager.subscribe_stream(id) { + Ok(stream) => { + let _ = channel.send(stream); + } + Err(error) => error!(%error, "subscribe to recording stream"), } - }, + } RecordingManagerMessage::ListFiles { id, channel } => { match manager.ongoing_recordings.get(&id) { Some(recording) => { diff --git a/devolutions-gateway/src/streaming.intent.md b/devolutions-gateway/src/streaming.intent.md index f3be7928e..14f95a120 100644 --- a/devolutions-gateway/src/streaming.intent.md +++ b/devolutions-gateway/src/streaming.intent.md @@ -46,4 +46,70 @@ Only WebM, asciicast, and TRP recording artifacts are accepted by the `/shadow` JREC artifact handling, storage, download content types, and consumer-side rendering are outside the scope of this document. -> **Boundary:** Session Recording Log artifacts are supported elsewhere in Gateway through the JREC recording flow. Their rejection by `/shadow` applies only to the WebSocket streaming path covered by this document. \ No newline at end of file +> **Boundary:** Session Recording Log artifacts are supported elsewhere in Gateway through the JREC recording flow. Their rejection by `/shadow` applies only to the WebSocket streaming path covered by this document. + + +## Multi-clip, size variant streaming + +We would like to support streamings of multi-clip, size-variant source. +See how we do recordings in `devolutions-gateway/src/recording.rs`. We now would like to support streaming as well for the same source. + +### The source +We have two streaming sources that we currently support: +1. RDM, which whenver size of a remote connecti session changes, it creates a new clip with consistent size in the header. +2. Chrome/Other browsers, chrome behaves differently, see `webapp/packages/web-recorder`, we use the media recorder API to record the session, the size changing behavior is not documented, but in experiencemnt and in practice, it will sliently change the size of the frame, the webm standard did not advise against this behavior, more lilely, it is undifined, and the client may or may not support it. + +### The normalizer + +Given the constrains above, we would like to unifiy the source and provide a single shape that the client can consume easily without breaking backward compatibility. +We would use the following model: + +```text +Legend: +----- size A +===== size B +^^^^^ size C +| input clip boundary ++ client joins +[ ] normalized output segment + +Time ------------------------------------------------------------------> + +RDM source: each size change creates a new input clip + +Input: ----------------------|======================|^^^^^^^^^^^^^^^^ + <------ clip 1 ------> <------ clip 2 ------> <--- clip 3 ---> + +Client 1: +[-----------][======================][^^^^^^^^^^^^^^^^] + starts near + end of clip 1 + +Client 2: +[==============][^^^^^^^^^^^^^^^^] + starts partway + through clip 2 + + +Browser source: sizes change inside one input clip + +Input: -----------------------=======================^^^^^^^^^^^^^^^^^ + <------------------- one input clip --------------------------> + +Client 1: +[------------][======================][^^^^^^^^^^^^^^^^] + starts near + end of size A + +Client 2: +[===============][^^^^^^^^^^^^^^^^] + starts partway + through size B + + +Normalized client output: + +- Each client begins at its own live edge. +- Each `[segment]` contains one fixed frame size. +- RDM input clip boundaries and browser frame-size changes produce the same + normalized output shape. +- Every client has an independent output sequence beginning at zero. +``` + +The client always gets a guaranteed fixed size segment, for the first segment, we keep it backward compatible, the protocol will be extendned, such that, on new `pull` message, when the previous output segment ends, it will send a new `SegmentStarted` message. \ No newline at end of file diff --git a/devolutions-gateway/src/streaming.rs b/devolutions-gateway/src/streaming.rs index 61dca69f2..4de037d5c 100644 --- a/devolutions-gateway/src/streaming.rs +++ b/devolutions-gateway/src/streaming.rs @@ -1,3 +1,5 @@ +use std::future::Future; +use std::pin::Pin; use std::sync::Arc; use std::time::Duration; @@ -5,36 +7,45 @@ use anyhow::Context; use axum::body::Body; use axum::extract::ws::{CloseFrame, Utf8Bytes, WebSocket}; use axum::response::Response; -use futures::SinkExt; +use devolutions_gateway_task::{ChildTask, ShutdownSignal}; +use futures::{SinkExt, Stream, stream}; use terminal_streamer::terminal_stream; -use tokio::fs::OpenOptions; -use tokio::sync::Notify; +use tokio::fs::{File, OpenOptions}; +use tokio::sync::{Notify, watch}; use uuid::Uuid; -use video_streamer::config::CpuCount; -use video_streamer::{ReOpenableFile, webm_stream}; +use video_streamer::{RecordingClip, RecordingEvent, RecordingSource, SessionConfig, StartAt, stream_session}; +use crate::recording::{RecordingMessageSender, RecordingStreamState}; use crate::token::RecordingFileType; -pub(crate) async fn stream_file( - path: &camino::Utf8Path, +pub(crate) async fn stream_recording( ws: axum::extract::WebSocketUpgrade, - shutdown_notify: Arc, - recordings: crate::recording::RecordingMessageSender, + shutdown_signal: ShutdownSignal, + recordings: RecordingMessageSender, recording_id: Uuid, ) -> anyhow::Result> { - let streaming_type = validate_streaming_file(path).await?; - - let when_new_chunk_appended = move || { - let (tx, rx) = tokio::sync::oneshot::channel(); - recordings.add_new_chunk_listener(recording_id, tx); - rx + let stream_state = recordings.subscribe_to_stream(recording_id).await?; + let (path, clip_sequence) = { + let state = stream_state.borrow(); + let clip = state.clips.last().context("recording has no clips")?; + (clip.path.clone(), clip.sequence) }; - - let path = Arc::new(path.to_owned()); + let streaming_type = validate_streaming_file(&path).await?; let upgrade_result = match streaming_type { StreamingType::Terminal(input_type) => { - let shutdown_notify = Arc::clone(&shutdown_notify); + let when_new_chunk_appended = move || { + let (tx, rx) = tokio::sync::oneshot::channel(); + recordings.add_new_chunk_listener(recording_id, tx); + rx + }; + let path = Arc::new(path); ws.on_upgrade(move |socket| async move { + let shutdown_notify = Arc::new(Notify::new()); + let notify = Arc::clone(&shutdown_notify); + let _shutdown_bridge = ChildTask::spawn(async move { + wait_for_terminal_stream_end(stream_state, clip_sequence, shutdown_signal).await; + notify.notify_one(); + }); if let Err(e) = setup_terminal_streaming(&path, input_type, socket, shutdown_notify, when_new_chunk_appended).await { @@ -42,14 +53,11 @@ pub(crate) async fn stream_file( } }) } - StreamingType::WebM => { - let shutdown_notify = Arc::clone(&shutdown_notify); - ws.on_upgrade(move |socket| async move { - if let Err(e) = setup_webm_streaming(&path, socket, shutdown_notify, when_new_chunk_appended).await { - error!(error = ?e, "WebM streaming failed"); - } - }) - } + StreamingType::WebM => ws.on_upgrade(move |socket| async move { + if let Err(e) = setup_webm_streaming(stream_state, socket, shutdown_signal).await { + error!(error = ?e, "WebM streaming failed"); + } + }), }; Ok(upgrade_result) @@ -142,44 +150,56 @@ async fn setup_terminal_streaming( Ok(()) } +async fn wait_for_recording_clip_end(mut stream_state: watch::Receiver, clip_sequence: u64) { + loop { + if stream_state + .borrow() + .active + .is_none_or(|active| active.sequence != clip_sequence) + { + return; + } + if stream_state.changed().await.is_err() { + return; + } + } +} + +async fn wait_for_terminal_stream_end( + stream_state: watch::Receiver, + clip_sequence: u64, + mut shutdown_signal: ShutdownSignal, +) { + tokio::select! { + () = wait_for_recording_clip_end(stream_state, clip_sequence) => {} + () = shutdown_signal.wait() => {} + } +} + async fn setup_webm_streaming( - path: &camino::Utf8Path, + stream_state: watch::Receiver, socket: WebSocket, - shutdown_notify: Arc, - when_new_chunk_appended: impl Fn() -> tokio::sync::oneshot::Receiver<()> + Send + 'static, + shutdown_signal: ShutdownSignal, ) -> anyhow::Result<()> { - let streaming_file = ReOpenableFile::open(path).with_context(|| format!("failed to open file: {path:?}"))?; - let streamer_config = video_streamer::StreamingConfig { - encoder_threads: CpuCount::default(), - adaptive_frame_skip: true, + let source = WebmRecordingSource { stream_state }; + let mut session_shutdown = shutdown_signal.clone(); + let (websocket_stream, close_handle) = crate::ws::handle_messages( + socket, + crate::ws::KeepAliveShutdownSignal(shutdown_signal), + Duration::from_secs(45), + ); + let streaming_result = tokio::select! { + result = stream_session(source, websocket_stream, SessionConfig::default()) => result, + () = session_shutdown.wait() => return Ok(()), }; - let (websocket_stream, close_handle) = - crate::ws::handle(socket, Arc::clone(&shutdown_notify), Duration::from_secs(45)); - let streaming_result = tokio::task::spawn_blocking(move || { - webm_stream( - websocket_stream, - streaming_file, - shutdown_notify, - streamer_config, - when_new_chunk_appended, - ) - .context("webm_stream failed")?; - Ok::<_, anyhow::Error>(()) - }) - .await; - match streaming_result { - Err(e) => { - error!(error=?e, "Streaming file task join failed"); - Err(anyhow::anyhow!("Streaming task failed")) - } - Ok(Err(e)) => { + Err(error) => { close_handle.server_error("webm streaming failure".to_owned()).await; - error!(error = format!("{e:#}"), "Streaming file failed"); - Err(e) + error!(error = format!("{error:#}"), "WebM streaming failed"); + Err(error) } - Ok(Ok(())) => { + Ok(()) => { close_handle.normal_close().await; Ok(()) } @@ -187,7 +207,7 @@ async fn setup_webm_streaming( } #[cfg(test)] -mod tests { +mod file_type_tests { use super::*; #[tokio::test] @@ -253,3 +273,351 @@ mod tests { assert!(streaming_type_for_file_type(RecordingFileType::SessionRecordingLog).is_err()); } } + +struct WebmRecordingSource { + stream_state: watch::Receiver, +} + +impl RecordingSource for WebmRecordingSource { + type Stream = Pin> + Send>>; + type Start = Pin> + Send>>; + + fn start(self) -> Self::Start { + Box::pin(async move { recording_event_stream(self.stream_state) }) + } +} + +struct CurrentRecordingClip { + sequence: u64, + caught_up: bool, +} + +struct RecordingEventSource { + stream_state: watch::Receiver, + next_clip: usize, + current_clip: Option, + next_start_at: StartAt, + ended: bool, +} + +impl RecordingEventSource { + fn new(mut stream_state: watch::Receiver) -> anyhow::Result { + let state = stream_state.borrow_and_update().clone(); + let (next_clip, next_start_at) = match state.active { + Some(active) => ( + usize::try_from(active.sequence).context("recording sequence does not fit in usize")?, + if active.ready { + StartAt::LiveEdge + } else { + StartAt::Beginning + }, + ), + None => (state.clips.len(), StartAt::Beginning), + }; + + Ok(Self { + stream_state, + next_clip, + current_clip: None, + next_start_at, + ended: false, + }) + } + + async fn next_event(&mut self) -> anyhow::Result> { + if self.ended { + return Ok(None); + } + + loop { + let state = self.stream_state.borrow().clone(); + + if let Some(current_clip) = self.current_clip.as_mut() { + if !current_clip.caught_up { + current_clip.caught_up = true; + return Ok(Some(RecordingEvent::CaughtUp)); + } + + if state + .active + .is_some_and(|active| active.sequence == current_clip.sequence) + { + if self.stream_state.has_changed()? { + let latest = self.stream_state.borrow_and_update().clone(); + if latest + .active + .is_some_and(|active| active.sequence == current_clip.sequence) + { + return Ok(Some(RecordingEvent::DataAvailable)); + } + continue; + } + self.stream_state + .changed() + .await + .context("recording stream state closed")?; + if self + .stream_state + .borrow() + .active + .is_some_and(|active| active.sequence == current_clip.sequence) + { + return Ok(Some(RecordingEvent::DataAvailable)); + } + continue; + } + + self.current_clip = None; + self.next_clip = self.next_clip.checked_add(1).context("recording clip index overflow")?; + return Ok(Some(RecordingEvent::ClipEnded)); + } + + if let Some(clip) = state.clips.get(self.next_clip) { + let expected_sequence = + u64::try_from(self.next_clip).context("recording clip index does not fit in u64")?; + if clip.sequence != expected_sequence { + anyhow::bail!("recording clip sequence is not contiguous"); + } + + if state + .active + .is_some_and(|active| active.sequence == clip.sequence && !active.ready) + { + self.stream_state + .changed() + .await + .context("recording stream state closed")?; + continue; + } + + if clip.path.extension() != Some(RecordingFileType::WebM.extension()) { + anyhow::bail!("recording clip is not WebM"); + } + + let file = File::open(&clip.path) + .await + .with_context(|| format!("failed to open recording clip: {}", clip.path))?; + let file = file.into_std().await; + let start_at = std::mem::replace(&mut self.next_start_at, StartAt::Beginning); + self.current_clip = Some(CurrentRecordingClip { + sequence: clip.sequence, + caught_up: false, + }); + return Ok(Some(RecordingEvent::ClipStarted { + sequence: clip.sequence, + start_at, + clip: RecordingClip::new(file), + })); + } + + if state.ended { + self.ended = true; + return Ok(Some(RecordingEvent::SessionEnded)); + } + + self.stream_state + .changed() + .await + .context("recording stream state closed")?; + } + } +} + +fn recording_event_stream( + stream_state: watch::Receiver, +) -> anyhow::Result> + Send>>> { + let source = RecordingEventSource::new(stream_state)?; + Ok(Box::pin(stream::unfold(Some(source), |source| async move { + let mut source = source?; + match source.next_event().await { + Ok(Some(event)) => Some((Ok(event), Some(source))), + Ok(None) => None, + Err(error) => Some((Err(error), None)), + } + }))) +} + +#[cfg(test)] +mod tests { + use std::fs; + + use super::*; + use crate::recording::{ActiveRecordingStreamClip, RecordingStreamClip}; + + struct ScratchDirectory(camino::Utf8PathBuf); + + impl Drop for ScratchDirectory { + fn drop(&mut self) { + let _ = fs::remove_dir_all(&self.0); + } + } + + #[tokio::test] + async fn terminal_clip_end_is_retained_before_waiting() { + let state = RecordingStreamState::for_test( + Vec::new(), + Some(ActiveRecordingStreamClip { + sequence: 0, + ready: true, + }), + false, + ); + let (sender, receiver) = watch::channel(state); + sender.send_modify(RecordingStreamState::mark_disconnected); + + tokio::time::timeout(Duration::from_millis(25), wait_for_recording_clip_end(receiver, 0)) + .await + .expect("clip end should already be visible"); + } + + #[tokio::test] + async fn reconnect_waits_for_the_next_clip_before_ending_the_session() { + let scratch = camino::Utf8PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .join("..") + .join("target") + .join("streaming-tests") + .join(Uuid::new_v4().to_string()); + fs::create_dir_all(&scratch).expect("create test directory"); + let _cleanup = ScratchDirectory(scratch.clone()); + + let first_path = scratch.join("recording-0.webm"); + let second_path = scratch.join("recording-1.webm"); + fs::write(&first_path, b"first").expect("write first clip"); + fs::write(&second_path, b"second").expect("write second clip"); + + let first_clip = RecordingStreamClip { + sequence: 0, + path: first_path, + }; + let state = RecordingStreamState::for_test( + vec![first_clip], + Some(ActiveRecordingStreamClip { + sequence: 0, + ready: true, + }), + false, + ); + let (sender, receiver) = watch::channel(state); + let mut source = RecordingEventSource::new(receiver).expect("create recording event source"); + + match source.next_event().await.expect("read first start") { + Some(RecordingEvent::ClipStarted { + sequence, + start_at, + clip, + }) => { + assert_eq!(sequence, 0); + assert_eq!(start_at, StartAt::LiveEdge); + drop(clip); + } + event => panic!("unexpected first start event: {event:?}"), + } + assert!(matches!( + source.next_event().await.expect("catch up first clip"), + Some(RecordingEvent::CaughtUp) + )); + sender.send_modify(|_| {}); + assert!(matches!( + source.next_event().await.expect("read availability marker"), + Some(RecordingEvent::DataAvailable) + )); + + sender.send_modify(RecordingStreamState::mark_disconnected); + assert!(matches!( + source.next_event().await.expect("end first clip"), + Some(RecordingEvent::ClipEnded) + )); + assert!( + tokio::time::timeout(Duration::from_millis(25), source.next_event()) + .await + .is_err(), + "a reconnectable disconnect must not emit SessionEnded" + ); + + sender.send_modify(|state| { + Arc::make_mut(&mut state.clips).push(RecordingStreamClip { + sequence: 1, + path: second_path, + }); + state.active = Some(ActiveRecordingStreamClip { + sequence: 1, + ready: true, + }); + state.ended = false; + }); + match source.next_event().await.expect("read second start") { + Some(RecordingEvent::ClipStarted { + sequence, + start_at, + clip, + }) => { + assert_eq!(sequence, 1); + assert_eq!(start_at, StartAt::Beginning); + drop(clip); + } + event => panic!("unexpected second start event: {event:?}"), + } + assert!(matches!( + source.next_event().await.expect("catch up second clip"), + Some(RecordingEvent::CaughtUp) + )); + + sender.send_modify(RecordingStreamState::mark_disconnected); + assert!(matches!( + source.next_event().await.expect("end second clip"), + Some(RecordingEvent::ClipEnded) + )); + sender.send_modify(RecordingStreamState::mark_ended); + assert!(matches!( + source.next_event().await.expect("end session"), + Some(RecordingEvent::SessionEnded) + )); + assert!(source.next_event().await.expect("finish source").is_none()); + } + + #[tokio::test] + async fn catch_up_precedes_coalesced_append_markers() { + let scratch = camino::Utf8PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .join("..") + .join("target") + .join("streaming-tests") + .join(Uuid::new_v4().to_string()); + fs::create_dir_all(&scratch).expect("create test directory"); + let _cleanup = ScratchDirectory(scratch.clone()); + + let path = scratch.join("recording-0.webm"); + fs::write(&path, b"recording").expect("write clip"); + let state = RecordingStreamState::for_test( + vec![RecordingStreamClip { sequence: 0, path }], + Some(ActiveRecordingStreamClip { + sequence: 0, + ready: true, + }), + false, + ); + let (sender, receiver) = watch::channel(state); + let mut source = RecordingEventSource::new(receiver).expect("create recording event source"); + + assert!(matches!( + source.next_event().await.expect("read clip start"), + Some(RecordingEvent::ClipStarted { .. }) + )); + sender.send_modify(|_| {}); + sender.send_modify(|_| {}); + + assert!(matches!( + source.next_event().await.expect("read catch-up marker"), + Some(RecordingEvent::CaughtUp) + )); + assert!(matches!( + source.next_event().await.expect("read coalesced availability marker"), + Some(RecordingEvent::DataAvailable) + )); + assert!( + tokio::time::timeout(Duration::from_millis(25), source.next_event()) + .await + .is_err(), + "coalesced append markers must produce one availability event" + ); + } +} diff --git a/devolutions-gateway/src/ws.rs b/devolutions-gateway/src/ws.rs index 59b667df6..a0ab14c6b 100644 --- a/devolutions-gateway/src/ws.rs +++ b/devolutions-gateway/src/ws.rs @@ -43,6 +43,51 @@ pub fn handle( (websocket_compat(ws), close_handle) } +pub fn handle_messages( + ws: WebSocket, + shutdown_signal: impl transport::KeepAliveShutdown, + keep_alive_interval: time::Duration, +) -> ( + impl futures::Stream> + + futures::Sink + + Unpin + + Send + + 'static, + transport::CloseWebSocketHandle, +) { + let ws = transport::Shared::new(ws); + + let close_handle = transport::spawn_websocket_sentinel_task( + ws.shared().with(|message: transport::WsWriteMsg| { + future::ready(Result::<_, axum::Error>::Ok(match message { + transport::WsWriteMsg::Ping => ws::Message::Ping(Bytes::new()), + transport::WsWriteMsg::Close(frame) => ws::Message::Close(Some(CloseFrame { + code: frame.code, + reason: frame.message.into(), + })), + })) + }), + shutdown_signal, + keep_alive_interval, + ); + + let messages = ws + .take_while(|item| future::ready(!matches!(item, Ok(ws::Message::Close(_))))) + .filter_map(|item| { + item.map(|msg| match msg { + ws::Message::Text(s) => Some(Bytes::from(s)), + ws::Message::Binary(data) => Some(data), + ws::Message::Ping(_) | ws::Message::Pong(_) => None, + ws::Message::Close(_) => None, + }) + .transpose() + .pipe(future::ready) + }) + .with(|item: Bytes| futures::future::ready(Ok::<_, axum::Error>(ws::Message::Binary(item)))); + + (messages, close_handle) +} + fn websocket_compat(ws: transport::Shared) -> impl AsyncRead + AsyncWrite + Unpin + Send + 'static { let ws_compat = ws .filter_map(|item| {