Skip to main content

flowly_mpegts/
muxer.rs

1use std::{io::Write, pin::pin};
2
3use bytes::{Buf, BufMut, Bytes, BytesMut};
4use flowly::{Fourcc, Frame, Service};
5use futures::StreamExt;
6use mpeg2ts::{
7    Error as TsError,
8    es::{StreamId, StreamType},
9    pes::PesHeader,
10    time::{ClockReference, Timestamp},
11    ts::{
12        AdaptationField, ContinuityCounter, EsInfo, Pid, ProgramAssociation,
13        TransportScramblingControl, TsHeader, TsPacket, TsPacketWriter, TsPayload, VersionNumber,
14        WriteTsPacket,
15        payload::{self, Pat, Pmt},
16    },
17};
18
19use crate::Error;
20
21const PMT_PID: u16 = 256;
22const VIDEO_ES_PID: u16 = 257;
23// const AUDIO_ES_PID: u16 = 258;
24const PES_VIDEO_STREAM_ID: u8 = 224;
25// const PES_AUDIO_STREAM_ID: u8 = 192;
26
27#[derive(Default)]
28pub struct Mpeg2TsMuxer {
29    video_continuity_counter: ContinuityCounter,
30    header_sent: bool,
31    params_updated: bool,
32}
33
34impl Mpeg2TsMuxer {
35    fn write_packet(
36        &mut self,
37        writer: &mut TsPacketWriter<impl Write>,
38        ts: Timestamp,
39        unit: &[u8],
40        is_keyframe: bool,
41    ) -> Result<(), Error> {
42        let mut header = Self::default_ts_header(VIDEO_ES_PID, self.video_continuity_counter)?;
43        let mut annexb = Vec::with_capacity(unit.len() + 3);
44        annexb.extend_from_slice(&[0, 0, 1]);
45        annexb.extend_from_slice(unit);
46        let mut buf = &annexb[..];
47
48        let packet = {
49            let data = payload::Bytes::new(&buf.chunk()[..buf.remaining().min(150)])?;
50            buf.advance(data.len());
51
52            TsPacket {
53                header: header.clone(),
54                adaptation_field: is_keyframe.then(|| AdaptationField {
55                    discontinuity_indicator: false,
56                    random_access_indicator: true,
57                    es_priority_indicator: false,
58                    pcr: Some(ClockReference::from(ts)),
59                    opcr: None,
60                    splice_countdown: None,
61                    transport_private_data: Vec::new(),
62                    extension: None,
63                }),
64                payload: Some(TsPayload::Pes(payload::Pes {
65                    header: PesHeader {
66                        stream_id: StreamId::new(PES_VIDEO_STREAM_ID),
67                        priority: false,
68                        data_alignment_indicator: false,
69                        copyright: false,
70                        original_or_copy: false,
71                        pts: Some(ts),
72                        dts: None,
73                        escr: None,
74                    },
75                    pes_packet_len: 0,
76                    data,
77                })),
78            }
79        };
80
81        writer.write_ts_packet(&packet)?;
82        header.continuity_counter.increment();
83
84        while buf.has_remaining() {
85            let raw_payload =
86                payload::Bytes::new(&buf.chunk()[..buf.remaining().min(payload::Bytes::MAX_SIZE)])?;
87
88            buf.advance(raw_payload.len());
89
90            let packet = TsPacket {
91                header: header.clone(),
92                adaptation_field: None,
93                payload: Some(TsPayload::Raw(raw_payload)),
94            };
95
96            writer.write_ts_packet(&packet)?;
97            header.continuity_counter.increment();
98        }
99
100        self.video_continuity_counter = header.continuity_counter;
101        Ok(())
102    }
103}
104
105impl Mpeg2TsMuxer {
106    #[inline]
107    fn write_header<W: WriteTsPacket>(
108        &mut self,
109        writer: &mut W,
110        stream_type: StreamType,
111    ) -> Result<(), TsError> {
112        self.write_packets(
113            writer,
114            [
115                &Self::default_pat_packet(),
116                &Self::default_pmt_packet(stream_type),
117            ],
118        )?;
119
120        Ok(())
121    }
122
123    #[inline]
124    fn write_packets<'a, W: WriteTsPacket, P: IntoIterator<Item = &'a TsPacket>>(
125        &mut self,
126        writer: &mut W,
127        packets: P,
128    ) -> Result<(), TsError> {
129        packets
130            .into_iter()
131            .try_for_each(|pak| writer.write_ts_packet(pak))?;
132
133        Ok(())
134    }
135
136    fn default_ts_header(
137        pid: u16,
138        continuity_counter: ContinuityCounter,
139    ) -> Result<TsHeader, TsError> {
140        Ok(TsHeader {
141            transport_error_indicator: false,
142            transport_priority: false,
143            pid: Pid::new(pid)?,
144            transport_scrambling_control: TransportScramblingControl::NotScrambled,
145            continuity_counter,
146        })
147    }
148
149    fn default_pat_packet() -> TsPacket {
150        TsPacket {
151            header: Self::default_ts_header(0, Default::default()).unwrap(),
152            adaptation_field: None,
153            payload: Some(TsPayload::Pat(Pat {
154                transport_stream_id: 1,
155                version_number: VersionNumber::default(),
156                table: vec![ProgramAssociation {
157                    program_num: 1,
158                    program_map_pid: Pid::new(PMT_PID).unwrap(),
159                }],
160            })),
161        }
162    }
163
164    fn default_pmt_packet(stream_type: StreamType) -> TsPacket {
165        TsPacket {
166            header: Self::default_ts_header(PMT_PID, Default::default()).unwrap(),
167            adaptation_field: None,
168            payload: Some(TsPayload::Pmt(Pmt {
169                program_num: 1,
170                pcr_pid: Some(Pid::new(VIDEO_ES_PID).unwrap()),
171                version_number: VersionNumber::default(),
172                program_info: vec![],
173                es_info: vec![EsInfo {
174                    stream_type,
175                    elementary_pid: Pid::new(VIDEO_ES_PID).unwrap(),
176                    descriptors: vec![],
177                }],
178            })),
179        }
180    }
181}
182
183impl Mpeg2TsMuxer {
184    fn push_frame<F: Frame>(&mut self, frame: F, dst: &mut BytesMut) -> Result<(), Error> {
185        let mut writer = TsPacketWriter::new(dst.writer());
186
187        if !self.header_sent {
188            self.header_sent = true;
189            self.write_header(
190                &mut writer,
191                match frame.codec() {
192                    Fourcc::VIDEO_AVC => StreamType::H264,
193                    Fourcc::VIDEO_HEVC => StreamType::H265,
194                    codec => return Err(Error::MuxUnsupportedCodec(codec)),
195                },
196            )?;
197        }
198
199        if frame.has_params() {
200            self.params_updated = true;
201        }
202
203        let ts = Timestamp::new((frame.pts() as u64 * 9) / 100).map_err(TsError::from)?;
204
205        for param in frame.params() {
206            if self.params_updated {
207                self.params_updated = false;
208            }
209
210            self.write_packet(&mut writer, ts, param, false)?;
211        }
212
213        for unit in frame.units() {
214            self.write_packet(&mut writer, ts, unit, frame.is_keyframe())?;
215        }
216
217        Ok(())
218    }
219}
220
221impl<F: Frame + Send, E: std::error::Error + Send + Sync + 'static> Service<Result<F, E>>
222    for Mpeg2TsMuxer
223{
224    type Out = Result<Bytes, Error<E>>;
225
226    fn handle(
227        mut self,
228        input: impl futures::Stream<Item = Result<F, E>> + Send,
229    ) -> impl futures::Stream<Item = Self::Out> + Send {
230        async_stream::stream! {
231            let mut input = pin!(input);
232            let mut buffer = BytesMut::new();
233
234            while let Some(res) = input.next().await {
235                match res {
236                    Ok(frame) => {
237                        if let Err(err) = self.push_frame(frame, &mut buffer) {
238                             yield Err(err.extend());
239                        }
240
241                        yield Ok(buffer.split().freeze());
242                    },
243                    Err(err) => yield Err(Error::Other(err)),
244                }
245            }
246        }
247    }
248}