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;
23const PES_VIDEO_STREAM_ID: u8 = 224;
25#[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}