diff --git a/crates/video-streamer/README.md b/crates/video-streamer/README.md index 1151304f2..31826d8d4 100644 --- a/crates/video-streamer/README.md +++ b/crates/video-streamer/README.md @@ -1,8 +1,18 @@ # video-streamer -This crate takes an unseekable WebM recording (typically from Chrome CaptureStream) and rewrites it into a “fresh” WebM stream that can start playing immediately. -It does this by parsing the incoming WebM, finding the correct cut point, and re-encoding frames. -The output stream begins with a keyframe and valid headers. +This crate takes an unseekable WebM recording and rewrites it into a stream that can start playing immediately. + +`webm_stream` still serves one growing file over the original Start/Pull protocol. +`stream_session` accepts a multi-clip recording event stream, reconnects across clips, and emits independent VP8 WebM segments over the same Start/Pull codes. + +The input event grammar is: + +```text +(ClipStarted Bytes* CaughtUp Bytes* ClipEnded)* SessionEnded +``` + +Pulls that arrive while a response is pending are queued. +`Stream ended` (type code 3) ends the session. ## Prerequisites 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..04e480466 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::{RecordingEvent, SessionConfig, StartAt, stream_session}; #[cfg(feature = "bench")] pub mod bench_support; diff --git a/crates/video-streamer/src/normalizer.rs b/crates/video-streamer/src/normalizer.rs new file mode 100644 index 000000000..1630bde5c --- /dev/null +++ b/crates/video-streamer/src/normalizer.rs @@ -0,0 +1,639 @@ +use std::io::{self, Write}; +use std::pin::Pin; +use std::task::{Context as TaskContext, Poll}; + +use anyhow::Context; +use bytes::{Bytes, BytesMut}; +use cadeau::xmf::vpx::{VpxCodec, VpxEncoder, VpxEncoderPreset, VpxImage}; +use ebml_iterable::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::{RecordingEvent, SessionConfig, StartAt}; +use crate::streamer::block_tag::{VideoBlock, is_vpx_key_frame}; + +const OUTPUT_CHANNEL_CAPACITY: usize = 1; +const INPUT_CHANNEL_CAPACITY: usize = 1; +const INPUT_CHUNK_SIZE: usize = 64 * 1024; +const MAX_BUFFERED_TAG_BYTES: usize = 64 * 1024 * 1024; +const MAX_PENDING_GOP_BYTES: usize = 64 * 1024 * 1024; +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; + +#[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(); + } + } +} + +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 } => { + anyhow::ensure!( + matches!(phase, SessionPhase::AwaitClip), + "clip {sequence} started before the previous clip ended" + ); + phase = SessionPhase::InClip(Box::new(ClipNormalizer::new( + sequence, + start_at, + sender.clone(), + config, + next_segment_sequence, + ))); + } + RecordingEvent::Bytes(bytes) => { + let SessionPhase::InClip(clip) = &mut phase else { + anyhow::bail!("recording bytes arrived outside a clip"); + }; + clip.push(&bytes)?; + } + 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, +} + +struct PendingFrame { + data: Vec, + timestamp: u64, + codec: VpxCodec, + key_frame: bool, +} + +enum ClipPhase { + History(HistoryPolicy), + Live, +} + +enum HistoryPolicy { + EmitAll, + KeepLatestGop(PendingGop), +} + +#[derive(Default)] +struct PendingGop { + frames: Vec, + bytes: usize, +} + +impl PendingGop { + fn push(&mut self, frame: PendingFrame) -> anyhow::Result<()> { + if frame.key_frame { + self.frames.clear(); + self.bytes = 0; + } else if self.frames.is_empty() { + return Ok(()); + } + + let bytes = self + .bytes + .checked_add(frame.data.len()) + .context("pending GOP size overflow")?; + anyhow::ensure!(bytes <= MAX_PENDING_GOP_BYTES, "pending GOP exceeds the resource limit"); + self.frames.push(frame); + self.bytes = bytes; + Ok(()) + } +} + +struct ClipNormalizer { + clip_sequence: u64, + decoder: TagDecoder, + input: BytesMut, + source_video: Option, + cluster_timestamp: Option, + timestamp_scale_ns: u64, + phase: ClipPhase, + input_decoder: Option, + output_segment: Option, + next_segment_sequence: u64, + sender: mpsc::Sender>, + config: SessionConfig, +} + +impl ClipNormalizer { + fn new( + clip_sequence: u64, + start_at: StartAt, + sender: mpsc::Sender>, + config: SessionConfig, + next_segment_sequence: u64, + ) -> Self { + let targets = [ + MatroskaSpec::TrackEntry(Master::Start), + MatroskaSpec::BlockGroup(Master::Start), + ]; + let mut decoder = TagDecoder::new(&targets); + decoder.set_max_allowable_tag_size(Some(MAX_BUFFERED_TAG_BYTES)); + let phase = match start_at { + StartAt::Beginning => ClipPhase::History(HistoryPolicy::EmitAll), + StartAt::LiveEdge => ClipPhase::History(HistoryPolicy::KeepLatestGop(PendingGop::default())), + }; + Self { + clip_sequence, + decoder, + input: BytesMut::new(), + source_video: None, + cluster_timestamp: None, + timestamp_scale_ns: WEBM_TIMESTAMP_SCALE_NS, + phase, + input_decoder: None, + output_segment: None, + next_segment_sequence, + sender, + config, + } + } + + fn push(&mut self, bytes: &[u8]) -> anyhow::Result<()> { + for chunk in bytes.chunks(INPUT_CHUNK_SIZE) { + self.input.extend_from_slice(chunk); + while let Some(positioned) = self.decoder.decode(&mut self.input)? { + self.handle_tag(positioned.tag)?; + } + } + Ok(()) + } + + fn caught_up(&mut self) -> 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 let HistoryPolicy::KeepLatestGop(pending) = history { + for frame in pending.frames { + self.process_frame(frame)?; + } + } + Ok(()) + } + + fn finish(mut self) -> anyhow::Result { + anyhow::ensure!( + matches!(self.phase, ClipPhase::Live), + "clip {} ended before caught-up", + self.clip_sequence + ); + loop { + match self.decoder.decode_eof(&mut self.input)? { + Some(positioned) => self.handle_tag(positioned.tag)?, + None if self.decoder.is_finished() => break, + None => continue, + } + } + + if let Some(segment) = self.output_segment.take() { + segment.finish()?; + } + Ok(self.next_segment_sequence) + } + + fn handle_tag(&mut self, tag: MatroskaSpec) -> anyhow::Result<()> { + match tag { + MatroskaSpec::TrackEntry(Master::Full(children)) => { + if let Some(video) = parse_video_track(&children)? { + anyhow::ensure!(self.source_video.is_none(), "multiple video tracks are not supported"); + self.source_video = Some(video); + } + } + MatroskaSpec::TimestampScale(value) => self.timestamp_scale_ns = value, + MatroskaSpec::Cluster(Master::Start) => self.cluster_timestamp = None, + MatroskaSpec::Timestamp(value) => self.cluster_timestamp = Some(value), + tag @ (MatroskaSpec::SimpleBlock(_) | MatroskaSpec::BlockGroup(Master::Full(_))) => { + self.handle_block(tag)?; + } + _ => {} + } + Ok(()) + } + + fn handle_block(&mut self, tag: MatroskaSpec) -> anyhow::Result<()> { + let video = self + .source_video + .context("video track header not found before video data")?; + let block = VideoBlock::new(tag, self.cluster_timestamp, video.codec)?; + if block.track != video.track { + return Ok(()); + } + + 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)?; + let frame = PendingFrame { + data, + timestamp, + codec: video.codec, + key_frame, + }; + + match &mut self.phase { + ClipPhase::History(HistoryPolicy::KeepLatestGop(pending)) => pending.push(frame), + ClipPhase::History(HistoryPolicy::EmitAll) | ClipPhase::Live => self.process_frame(frame), + } + } + + fn process_frame(&mut self, frame: PendingFrame) -> 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 size_changed = self + .output_segment + .as_ref() + .is_some_and(|current| current.dimensions != dimensions); + if size_changed { + self.output_segment + .take() + .context("missing active output segment")? + .finish()?; + } + + let new_segment = if self.output_segment.is_none() { + Some(SegmentInfo { + sequence: self.next_segment_sequence, + width: dimensions.width, + height: dimensions.height, + }) + } else { + None + }; + + 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 parse_video_track(children: &[MatroskaSpec]) -> anyhow::Result> { + let is_video = children + .iter() + .find_map(|tag| match tag { + MatroskaSpec::TrackType(value) => Some(*value == 1), + _ => None, + }) + .unwrap_or(false); + + if !is_video { + return Ok(None); + } + + let track = children + .iter() + .find_map(|tag| match tag { + MatroskaSpec::TrackNumber(value) => Some(*value), + _ => None, + }) + .context("video track number is missing")?; + let codec_id = children + .iter() + .find_map(|tag| match tag { + MatroskaSpec::CodecID(value) => Some(value.as_str()), + _ => None, + }) + .context("video codec ID is missing")?; + let codec = match codec_id { + "V_VP8" | "vp8" => VpxCodec::VP8, + "V_VP9" | "vp9" => VpxCodec::VP9, + _ => anyhow::bail!("unsupported video codec: {codec_id}"), + }; + + Ok(Some(SourceVideo { track, codec })) +} + +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 { + self.sender + .blocking_send(Ok(SegmentEvent::Data(Bytes::copy_from_slice(buffer)))) + .map_err(|_| io::Error::new(io::ErrorKind::BrokenPipe, "segment event receiver closed"))?; + Ok(buffer.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} diff --git a/crates/video-streamer/src/protocol.rs b/crates/video-streamer/src/protocol.rs new file mode 100644 index 000000000..ea375a241 --- /dev/null +++ b/crates/video-streamer/src/protocol.rs @@ -0,0 +1,440 @@ +use std::collections::VecDeque; +use std::error::Error; +use std::pin::Pin; + +use anyhow::Context as _; +use bytes::{BufMut as _, Bytes, BytesMut}; +use futures_util::{Sink, SinkExt as _, Stream, StreamExt as _}; + +use crate::normalizer::{SegmentEvent, SegmentInfo}; + +#[derive(Debug, Eq, PartialEq)] +pub(crate) enum ServerMessage { + Chunk(Bytes), + SegmentStarted(SegmentInfo), + Error(UserFriendlyError), + StreamEnded, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum ClientMessage { + Start, + Pull, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub(crate) enum UserFriendlyError { + UnexpectedError, +} + +impl UserFriendlyError { + fn as_str(&self) -> &'static str { + match self { + Self::UnexpectedError => "UnexpectedError", + } + } +} + +pub(crate) async fn stream_segments(mut transport: T, segments: S) -> anyhow::Result<()> +where + T: Stream> + Sink + Unpin, + S: Stream>, + E: Error + Send + Sync + 'static, +{ + tokio::pin!(segments); + let mut expected = ClientMessage::Start; + let mut segment_state = SegmentState::AwaitingBegin; + let mut queued_pulls = VecDeque::new(); + + loop { + let message = if let Some(message) = queued_pulls.pop_front() { + Ok(message) + } else { + let Some(message) = transport.next().await else { + return Ok(()); + }; + let message = message + .map_err(anyhow::Error::new) + .context("read client stream message")?; + decode_client_message(&message) + }; + let message = match message { + Ok(message) if message == expected => message, + Ok(message) => { + debug!( + expected = ?expected, + got = ?message, + "Rejected client request in wrong state" + ); + let _ = + send_server_message(&mut transport, ServerMessage::Error(UserFriendlyError::UnexpectedError)).await; + anyhow::bail!("invalid client stream state"); + } + Err(error) => { + debug!(error = %error, "Rejected undecodable client request"); + let _ = + send_server_message(&mut transport, ServerMessage::Error(UserFriendlyError::UnexpectedError)).await; + anyhow::bail!("invalid client stream state"); + } + }; + debug!( + request = ?message, + segment_state = ?segment_state, + queued_pulls = queued_pulls.len(), + "Serving client request" + ); + + let response = match wait_for_response(&mut transport, segments.as_mut(), &mut segment_state, &mut queued_pulls) + .await + { + Ok(Some(response)) => { + debug!( + response = ?response_kind(&response), + queued_pulls = queued_pulls.len(), + "Sending server response" + ); + response + } + Ok(None) => return Ok(()), + Err(error) => { + debug!(error = format!("{error:#}"), "Request failed while waiting"); + let _ = + send_server_message(&mut transport, ServerMessage::Error(UserFriendlyError::UnexpectedError)).await; + return Err(error); + } + }; + + let ended = response == ServerMessage::StreamEnded; + send_server_message(&mut transport, response).await?; + if ended { + return Ok(()); + } + + expected = match message { + ClientMessage::Start | ClientMessage::Pull => ClientMessage::Pull, + }; + } +} + +async fn send_server_message(transport: &mut T, message: ServerMessage) -> anyhow::Result<()> +where + T: Sink + Unpin, + E: Error + Send + Sync + 'static, +{ + transport + .send(encode_server_message(message)) + .await + .map_err(anyhow::Error::new) + .context("write server stream message") +} + +fn decode_client_message(message: &[u8]) -> anyhow::Result { + match message { + [0] => Ok(ClientMessage::Start), + [1] => Ok(ClientMessage::Pull), + _ => anyhow::bail!("invalid client message"), + } +} + +fn response_kind(message: &ServerMessage) -> &'static str { + match message { + ServerMessage::Chunk(_) => "chunk", + ServerMessage::SegmentStarted(_) => "segment-started", + ServerMessage::Error(_) => "error", + ServerMessage::StreamEnded => "stream-ended", + } +} + +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::SegmentStarted(info) => { + encoded.put_u8(1); + let json = format!( + "{{\"codec\":\"vp8\",\"sequence\":{},\"width\":{},\"height\":{}}}", + info.sequence, info.width, info.height + ); + encoded.put(json.as_bytes()); + } + 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), + } + encoded.freeze() +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum SegmentState { + AwaitingBegin, + Streaming, +} + +async fn wait_for_response( + transport: &mut T, + mut segments: Pin<&mut S>, + state: &mut SegmentState, + queued_pulls: &mut VecDeque, +) -> anyhow::Result> +where + T: Stream> + Unpin, + S: Stream>, + E: Error + Send + Sync + 'static, +{ + loop { + tokio::select! { + biased; + response = next_segment_message(segments.as_mut(), state) => { + return response; + } + message = transport.next() => match message { + None => return Ok(None), + Some(Ok(message)) => { + let message = decode_client_message(&message).context("decode pipelined client request")?; + if message == ClientMessage::Pull { + debug!( + queued_pulls = queued_pulls.len() + 1, + "Queued pipelined Pull while waiting for response" + ); + queued_pulls.push_back(message); + } else { + debug!( + incoming = ?message, + "Overlapping non-Pull request while waiting for response" + ); + anyhow::bail!("client sent another request before receiving a response"); + } + } + Some(Err(error)) => { + return Err(anyhow::Error::new(error).context("read client stream message")); + } + }, + } + } +} + +async fn next_segment_message( + mut segments: Pin<&mut S>, + state: &mut SegmentState, +) -> anyhow::Result> +where + S: Stream>, +{ + loop { + let Some(event) = segments.as_mut().next().await else { + anyhow::ensure!( + *state == SegmentState::AwaitingBegin, + "segment stream ended inside a segment" + ); + return Ok(Some(ServerMessage::StreamEnded)); + }; + + match event? { + SegmentEvent::Begin(info) => { + anyhow::ensure!( + *state == SegmentState::AwaitingBegin, + "segment began before the previous segment ended" + ); + *state = SegmentState::Streaming; + debug!( + sequence = info.sequence, + width = info.width, + height = info.height, + "Segment begin" + ); + return Ok(Some(ServerMessage::SegmentStarted(info))); + } + SegmentEvent::Data(data) => { + anyhow::ensure!( + *state == SegmentState::Streaming, + "segment data arrived outside a segment" + ); + debug!(bytes = data.len(), "Segment data"); + return Ok(Some(ServerMessage::Chunk(data))); + } + SegmentEvent::End => { + anyhow::ensure!(*state == SegmentState::Streaming, "segment ended outside a segment"); + *state = SegmentState::AwaitingBegin; + debug!("Segment end"); + } + } + } +} + +#[cfg(test)] +mod tests { + use std::collections::VecDeque; + + use futures_util::{StreamExt as _, stream}; + + use super::*; + + fn pending_after( + messages: impl IntoIterator>, + ) -> impl Stream> + Unpin { + stream::iter(messages).chain(stream::pending()) + } + + #[test] + fn protocol_codes_are_stable() { + assert_eq!( + encode_server_message(ServerMessage::SegmentStarted(SegmentInfo { + sequence: 7, + width: 1920, + height: 1080, + })), + Bytes::from_static(b"\x01{\"codec\":\"vp8\",\"sequence\":7,\"width\":1920,\"height\":1080}") + ); + 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"\x00\x01").is_err()); + assert!(decode_client_message(b"").is_err()); + } + + #[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 segments = stream::iter(events); + tokio::pin!(segments); + let mut state = SegmentState::AwaitingBegin; + + assert!(matches!( + next_segment_message(segments.as_mut(), &mut state) + .await + .expect("first begin"), + Some(ServerMessage::SegmentStarted(SegmentInfo { sequence: 0, .. })) + )); + assert_eq!( + next_segment_message(segments.as_mut(), &mut state) + .await + .expect("first data"), + Some(ServerMessage::Chunk(Bytes::from_static(b"first"))) + ); + assert!(matches!( + next_segment_message(segments.as_mut(), &mut state) + .await + .expect("second begin"), + Some(ServerMessage::SegmentStarted(SegmentInfo { sequence: 1, .. })) + )); + assert_eq!( + next_segment_message(segments.as_mut(), &mut state) + .await + .expect("second data"), + Some(ServerMessage::Chunk(Bytes::from_static(b"second"))) + ); + assert_eq!( + next_segment_message(segments.as_mut(), &mut state) + .await + .expect("stream end"), + Some(ServerMessage::StreamEnded) + ); + } + + #[tokio::test] + async fn start_response_buffers_one_pipelined_pull() { + let mut transport = pending_after([Ok::<_, std::io::Error>(Bytes::from_static(b"\x01"))]); + let segments = stream::once(async { + tokio::task::yield_now().await; + Ok(SegmentEvent::Begin(SegmentInfo { + sequence: 0, + width: 640, + height: 480, + })) + }); + tokio::pin!(segments); + let mut state = SegmentState::AwaitingBegin; + let mut queued_pulls = VecDeque::new(); + + let response = wait_for_response(&mut transport, segments.as_mut(), &mut state, &mut queued_pulls) + .await + .expect("wait for start response") + .expect("segment response"); + + assert!(matches!( + response, + ServerMessage::SegmentStarted(SegmentInfo { sequence: 0, .. }) + )); + assert_eq!(queued_pulls, VecDeque::from([ClientMessage::Pull])); + } + + #[tokio::test] + async fn extra_pull_while_waiting_for_chunk_is_queued() { + let mut transport = pending_after([Ok::<_, std::io::Error>(Bytes::from_static(b"\x01"))]); + let segments = stream::once(async { + tokio::task::yield_now().await; + Ok(SegmentEvent::Data(Bytes::from_static(b"chunk"))) + }); + tokio::pin!(segments); + let mut state = SegmentState::Streaming; + let mut queued_pulls = VecDeque::new(); + + let response = wait_for_response(&mut transport, segments.as_mut(), &mut state, &mut queued_pulls) + .await + .expect("wait for chunk") + .expect("chunk response"); + + assert_eq!(response, ServerMessage::Chunk(Bytes::from_static(b"chunk"))); + assert_eq!(queued_pulls, VecDeque::from([ClientMessage::Pull])); + } + + #[tokio::test] + async fn extra_start_while_waiting_is_still_rejected() { + let mut transport = pending_after([Ok::<_, std::io::Error>(Bytes::from_static(b"\x00"))]); + let segments = stream::once(async { + tokio::task::yield_now().await; + Ok(SegmentEvent::Data(Bytes::from_static(b"chunk"))) + }); + tokio::pin!(segments); + let mut state = SegmentState::Streaming; + let mut queued_pulls = VecDeque::new(); + + let error = wait_for_response(&mut transport, segments.as_mut(), &mut state, &mut queued_pulls) + .await + .expect_err("overlapping Start must fail"); + + assert!( + format!("{error:#}").contains("client sent another request before receiving a response"), + "{error:#}" + ); + } +} diff --git a/crates/video-streamer/src/session.rs b/crates/video-streamer/src/session.rs new file mode 100644 index 000000000..a8711b36f --- /dev/null +++ b/crates/video-streamer/src/session.rs @@ -0,0 +1,55 @@ +use std::error::Error; + +use bytes::Bytes; +use futures_util::{Sink, Stream}; + +/// A structural event from one append-only recording session. +/// +/// A clip starts, receives zero or more byte events, catches up exactly once, receives more bytes, +/// and ends before another clip starts. +/// The session ends only when no clip is active. +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum RecordingEvent { + ClipStarted { sequence: u64, start_at: StartAt }, + Bytes(Bytes), + 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, +} + +impl Default for SessionConfig { + fn default() -> Self { + Self { + encoder_threads: u32::try_from(num_cpus::get()).unwrap_or(1).max(1), + } + } +} + +/// Converts a recording session into fixed-size VP8 WebM segments over one pull-driven stream. +pub async fn stream_session(source: S, transport: T, config: SessionConfig) -> anyhow::Result<()> +where + S: Stream> + Send + 'static, + T: Stream> + Sink + Unpin, + E: Error + Send + Sync + 'static, +{ + let mut segments = crate::normalizer::normalize(source, config); + let stream_result = crate::protocol::stream_segments(transport, &mut segments).await; + let shutdown_result = segments.shutdown().await; + + match (stream_result, shutdown_result) { + (Err(error), _) => Err(error), + (Ok(()), Err(error)) => Err(error), + (Ok(()), Ok(())) => Ok(()), + } +} diff --git a/crates/video-streamer/src/streamer/block_tag.rs b/crates/video-streamer/src/streamer/block_tag.rs index 483dfc108..769657ed6 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,11 +84,13 @@ 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 @@ -120,7 +126,7 @@ impl VideoBlock { } }; - assert!(frame.len() == 1); + anyhow::ensure!(frame.len() == 1, "laced video blocks are not supported"); Ok(frame[0].clone()) } } 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 b010f13a6..18df871e4 100644 --- a/devolutions-gateway/src/api/jrec.rs +++ b/devolutions-gateway/src/api/jrec.rs @@ -949,7 +949,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, @@ -962,31 +966,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..b68bc338c 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::{Notify, 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,27 @@ 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, + revision: u64, } enum RecordingManagerMessage { @@ -253,6 +289,12 @@ enum RecordingManagerMessage { Disconnect { id: Uuid, }, + ClipStarted { + id: Uuid, + }, + ChunkAppended { + id: Uuid, + }, GetState { id: Uuid, channel: oneshot::Sender>, @@ -272,6 +314,10 @@ enum RecordingManagerMessage { id: Uuid, channel: oneshot::Sender>, }, + SubscribeToStream { + id: Uuid, + channel: oneshot::Sender>, + }, } impl fmt::Debug for RecordingManagerMessage { @@ -289,6 +335,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() } @@ -307,6 +355,10 @@ impl fmt::Debug for RecordingManagerMessage { 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,18 +438,28 @@ 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> { @@ -411,6 +473,20 @@ impl RecordingMessageSender { Ok(rx.await?) } + pub(crate) async fn subscribe_to_stream( + &self, + recording_id: Uuid, + ) -> anyhow::Result> { + let (tx, rx) = oneshot::channel(); + self.channel + .send(RecordingManagerMessage::SubscribeToStream { + id: recording_id, + channel: tx, + }) + .await?; + Ok(rx.await?) + } + pub(crate) async fn list_files(&self, recording_id: Uuid) -> anyhow::Result> { let (tx, rx) = oneshot::channel(); self.channel @@ -516,6 +592,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 +668,45 @@ 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; + state.revision = state.revision.saturating_add(1); + }); + 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, + revision: 0, + }; + watch::channel(state).0 + }; + self.ongoing_recordings.insert( id, OnGoingRecording { @@ -596,6 +715,7 @@ impl RecordingManagerTask { manifest_path, session_must_be_recorded, disconnected_ttl, + stream_state, }, ); let ongoing_recording_count = self.ongoing_recordings.len(); @@ -612,6 +732,54 @@ 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, + }); + state.revision = state.revision.saturating_add(1); + }); + + 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(|state| { + state.revision = state.revision.saturating_add(1); + }); + + 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,6 +815,12 @@ impl RecordingManagerTask { .save_to_file(&ongoing.manifest_path) .with_context(|| format!("write manifest at {}", ongoing.manifest_path))?; + ongoing.stream_state.send_modify(|state| { + state.active = None; + state.ended = true; + state.revision = state.revision.saturating_add(1); + }); + // Notify all the streamers that recording has ended. if let Some(notify) = self.recording_end_notifier.get(&id) { notify.notify_waiters(); @@ -686,6 +860,11 @@ 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(|state| { + state.active = None; + state.ended = true; + state.revision = state.revision.saturating_add(1); + }); // Check the recording policy of the associated session and kill it if necessary. if ongoing.session_must_be_recorded { @@ -745,6 +924,14 @@ impl RecordingManagerTask { 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()) + } } #[async_trait] @@ -822,6 +1009,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); @@ -847,6 +1044,14 @@ async fn recording_manager_task( 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.rs b/devolutions-gateway/src/streaming.rs index b62f4994c..4bc280cb4 100644 --- a/devolutions-gateway/src/streaming.rs +++ b/devolutions-gateway/src/streaming.rs @@ -5,35 +5,43 @@ use anyhow::Context; use axum::body::Body; use axum::extract::ws::{CloseFrame, Utf8Bytes, WebSocket}; use axum::response::Response; -use futures::SinkExt; +use bytes::Bytes; +use devolutions_gateway_task::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::io::AsyncReadExt; +use tokio::sync::{Notify, watch}; use uuid::Uuid; -use video_streamer::config::CpuCount; -use video_streamer::{ReOpenableFile, webm_stream}; +use video_streamer::{RecordingEvent, 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 path = Arc::new(path.to_owned()); + let stream_state = recordings.subscribe_to_stream(recording_id).await?; + let path = stream_state + .borrow() + .clips + .last() + .context("recording has no clips")? + .path + .clone(); + let streaming_type = validate_streaming_file(&path).await?; let upgrade_result = match streaming_type { StreamingType::Terminal => { - let shutdown_notify = Arc::clone(&shutdown_notify); + let shutdown_notify = recordings.subscribe_to_recording_finish(recording_id).await?; + 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 { if let Err(e) = setup_terminal_streaming(&path, socket, shutdown_notify, when_new_chunk_appended).await { @@ -41,14 +49,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) @@ -150,45 +155,168 @@ async fn setup_terminal_streaming( } 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 (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; + let source = recording_event_stream(stream_state)?; + let (websocket_stream, close_handle) = crate::ws::handle_messages( + socket, + crate::ws::KeepAliveShutdownSignal(shutdown_signal), + Duration::from_secs(45), + ); + let streaming_result = stream_session(source, websocket_stream, SessionConfig::default()).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(()) } } } + +struct CurrentRecordingClip { + sequence: u64, + file: File, + 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> { + const READ_BUFFER_SIZE: usize = 64 * 1024; + + if self.ended { + return Ok(None); + } + + loop { + let state = self.stream_state.borrow_and_update().clone(); + + if let Some(current_clip) = self.current_clip.as_mut() { + let mut bytes = vec![0; READ_BUFFER_SIZE]; + let read = current_clip.file.read(&mut bytes).await?; + if read > 0 { + bytes.truncate(read); + return Ok(Some(RecordingEvent::Bytes(Bytes::from(bytes)))); + } + + 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) + { + self.stream_state + .changed() + .await + .context("recording stream state closed")?; + 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 start_at = std::mem::replace(&mut self.next_start_at, StartAt::Beginning); + self.current_clip = Some(CurrentRecordingClip { + sequence: clip.sequence, + file, + caught_up: false, + }); + return Ok(Some(RecordingEvent::ClipStarted { + sequence: clip.sequence, + start_at, + })); + } + + 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 + 'static> { + let source = RecordingEventSource::new(stream_state)?; + Ok(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)), + } + })) +} 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| {