Skip to main content

moq_sub/
media.rs

1// SPDX-FileCopyrightText: 2024-2026 Cloudflare Inc., Luke Curley, Mike English and contributors
2// SPDX-FileCopyrightText: 2023-2024 Luke Curley and contributors
3// SPDX-License-Identifier: MIT OR Apache-2.0
4
5use std::{io::Cursor, sync::Arc};
6
7use anyhow::Context;
8use moq_transport::serve::{
9    SubgroupObjectReader, SubgroupReader, TrackReader, TrackReaderMode, Tracks, TracksReader,
10    TracksWriter,
11};
12use moq_transport::session::Subscriber;
13use mp4::ReadBox;
14use tokio::{
15    io::{AsyncReadExt, AsyncWrite, AsyncWriteExt},
16    sync::Mutex,
17    task::JoinSet,
18};
19use tracing::{debug, info, trace, warn};
20
21pub struct Media<O> {
22    subscriber: Subscriber,
23    broadcast: TracksReader,
24    tracks_writer: TracksWriter,
25    output: Arc<Mutex<O>>,
26    request_catalog: bool,
27}
28
29impl<O: AsyncWrite + Send + Unpin + 'static> Media<O> {
30    pub async fn new(
31        subscriber: Subscriber,
32        tracks: Tracks,
33        output: O,
34        request_catalog: bool,
35    ) -> anyhow::Result<Self> {
36        let (tracks_writer, _tracks_request, tracks_reader) = tracks.produce();
37        let broadcast = tracks_reader; // breadcrumb for navigating API name changes
38        Ok(Self {
39            subscriber,
40            broadcast,
41            tracks_writer,
42            output: Arc::new(Mutex::new(output)),
43            request_catalog,
44        })
45    }
46
47    pub async fn run(&mut self) -> anyhow::Result<()> {
48        let catalog = if self.request_catalog {
49            // The catalog track has no standardized name, but
50            // both moq-pub of moq-rs and gst-moq-pub uses ".catalog".
51            let buf = self.download_first_object(".catalog", "catalog").await?;
52            let s = std::str::from_utf8(&buf)?;
53            let c: moq_catalog::Root = serde_json::from_str(s)?;
54            info!("catalog: {c:#?}");
55            anyhow::ensure!(c.version == 1, "Unknown catalog version");
56            Some(c)
57        } else {
58            None
59        };
60        let moov = {
61            let init_track_name: &str = match catalog {
62                Some(ref c) => &c.tracks[0].init_track.clone().unwrap(),
63                None => "0.mp4",
64            };
65            let buf = self.download_first_object(init_track_name, "init").await?;
66            self.output.lock().await.write_all(&buf).await?;
67            let mut reader = Cursor::new(&buf);
68
69            let ftyp = read_atom(&mut reader).await?;
70            anyhow::ensure!(&ftyp[4..8] == b"ftyp", "expected ftyp atom");
71
72            let moov = read_atom(&mut reader).await?;
73            anyhow::ensure!(&moov[4..8] == b"moov", "expected moov atom");
74            let mut moov_reader = Cursor::new(&moov);
75            let moov_header = mp4::BoxHeader::read(&mut moov_reader)?;
76
77            mp4::MoovBox::read_box(&mut moov_reader, moov_header.size)?
78        };
79
80        let mut has_video = false;
81        let mut has_audio = false;
82        let mut tracks = vec![];
83        for (idx, trak) in moov.traks.into_iter().enumerate() {
84            let id = trak.tkhd.track_id;
85            let name: String = match catalog {
86                Some(ref c) => c.tracks[idx].name.clone(),
87                None => format!("{id}.m4s"),
88            };
89            info!("found track {name}");
90            let mut active = false;
91            if !has_video && trak.mdia.minf.stbl.stsd.avc1.is_some() {
92                active = true;
93                has_video = true;
94                info!("using {name} for video");
95            }
96            if !has_audio && trak.mdia.minf.stbl.stsd.mp4a.is_some() {
97                active = true;
98                has_audio = true;
99                info!("using {name} for audio");
100            }
101            if active {
102                let track = self
103                    .tracks_writer
104                    .create(&name)
105                    .context("failed to create track")?;
106
107                let mut subscriber = self.subscriber.clone();
108                tokio::task::spawn(async move {
109                    subscriber.subscribe(track).await.unwrap_or_else(|err| {
110                        warn!("failed to subscribe to track: {err:?}");
111                    });
112                });
113
114                tracks.push(
115                    self.broadcast
116                        .subscribe(self.broadcast.namespace.clone(), &name)
117                        .context("no track")?,
118                );
119            }
120        }
121
122        info!("playing {} tracks", tracks.len());
123        let mut tasks = JoinSet::new();
124        for track in tracks {
125            let out = self.output.clone();
126            tasks.spawn(async move {
127                let name = track.name.clone();
128                if let Err(err) = Self::recv_track(track, out).await {
129                    warn!("failed to play track {name}: {err:?}");
130                }
131            });
132        }
133        while tasks.join_next().await.is_some() {}
134        Ok(())
135    }
136
137    async fn download_first_object(
138        &mut self,
139        track_name: &str,
140        alias: &'static str,
141    ) -> anyhow::Result<Vec<u8>> {
142        let track = self
143            .tracks_writer
144            .create(track_name)
145            .context(format!("failed to create {alias} track"))?;
146
147        let mut subscriber = self.subscriber.clone();
148        tokio::task::spawn(async move {
149            subscriber.subscribe(track).await.unwrap_or_else(|err| {
150                warn!("failed to subscribe to {alias} track: {err:?}");
151            });
152        });
153
154        let track = self
155            .broadcast
156            .subscribe(self.broadcast.namespace.clone(), track_name)
157            .context(format!("no {alias} track"))?;
158        let mut group = match track.mode().await? {
159            TrackReaderMode::Subgroups(mut groups) => {
160                groups.next().await?.context(format!("no {alias} group"))?
161            }
162            _ => anyhow::bail!("expected {alias} segment"),
163        };
164
165        let object = group
166            .next()
167            .await?
168            .context(format!("no {alias} fragment"))?;
169        let buf = Self::recv_object(object).await?;
170        Ok(buf)
171    }
172
173    async fn recv_track(track: TrackReader, out: Arc<Mutex<O>>) -> anyhow::Result<()> {
174        let name = track.name.clone();
175        debug!("track {name}: start");
176        if let TrackReaderMode::Subgroups(mut groups) = track.mode().await? {
177            while let Some(group) = groups.next().await? {
178                let out = out.clone();
179                if let Err(err) = Self::recv_group(group, out).await {
180                    warn!("failed to receive group: {err:?}");
181                }
182            }
183        }
184        debug!("track {name}: finish");
185        Ok(())
186    }
187
188    async fn recv_group(mut group: SubgroupReader, out: Arc<Mutex<O>>) -> anyhow::Result<()> {
189        trace!("group={} start", group.group_id);
190        while let Some(object) = group.next().await? {
191            trace!(
192                "group={} fragment={} start",
193                group.group_id,
194                object.object_id
195            );
196            let out = out.clone();
197            let buf = Self::recv_object(object).await?;
198
199            out.lock().await.write_all(&buf).await?;
200        }
201
202        Ok(())
203    }
204
205    async fn recv_object(mut object: SubgroupObjectReader) -> anyhow::Result<Vec<u8>> {
206        let mut buf = Vec::with_capacity(object.size);
207        while let Some(chunk) = object.read().await? {
208            buf.extend_from_slice(&chunk);
209        }
210        Ok(buf)
211    }
212}
213
214// Read a full MP4 atom into a vector.
215async fn read_atom<R: AsyncReadExt + Unpin>(reader: &mut R) -> anyhow::Result<Vec<u8>> {
216    // Read the 8 bytes for the size + type
217    let mut buf = [0u8; 8];
218    reader.read_exact(&mut buf).await?;
219
220    // Convert the first 4 bytes into the size.
221    let size = u32::from_be_bytes(buf[0..4].try_into()?) as u64;
222
223    let mut raw = buf.to_vec();
224
225    let mut limit = match size {
226        // Runs until the end of the file.
227        0 => reader.take(u64::MAX),
228
229        // The next 8 bytes are the extended size to be used instead.
230        1 => {
231            reader.read_exact(&mut buf).await?;
232            let size_large = u64::from_be_bytes(buf);
233            anyhow::ensure!(
234                size_large >= 16,
235                "impossible extended box size: {}",
236                size_large
237            );
238
239            reader.take(size_large - 16)
240        }
241
242        2..=7 => {
243            anyhow::bail!("impossible box size: {}", size)
244        }
245
246        size => reader.take(size - 8),
247    };
248
249    // Append to the vector and return it.
250    let _read_bytes = limit.read_to_end(&mut raw).await?;
251
252    Ok(raw)
253}