use broadcast_common::traits::{Parse, Serialize};
use dvb_si::descriptors::DescriptorLoop;
use dvb_si::tables::pat::{PatEntry, PatSection};
use dvb_si::tables::pmt::{PmtSection, PmtStream, StreamType};
use mpeg_ts::mux::SectionPacketiser;
use ts_fix::{PidFilter, TsFix};
const PAT_PID: u16 = 0x0000;
const PMT1_PID: u16 = 0x0100;
const PMT2_PID: u16 = 0x0200;
const P1_PCR_PID: u16 = 0x0101;
const P1_VIDEO_PID: u16 = 0x0101;
const P1_AUDIO_PID: u16 = 0x0102;
const P2_PCR_PID: u16 = 0x0201;
const P2_VIDEO_PID: u16 = 0x0201;
const P2_AUDIO_PID: u16 = 0x0202;
const NULL_PID: u16 = 0x1FFF;
fn dummy_es_packet(pid: u16, cc: u8) -> [u8; 188] {
let mut pkt = [0u8; 188];
pkt[0] = 0x47; pkt[1] = ((pid >> 8) as u8) & 0x1F; pkt[2] = (pid & 0xFF) as u8;
pkt[3] = 0x10 | (cc & 0x0F); for (i, b) in pkt[4..].iter_mut().enumerate() {
*b = (i as u8).wrapping_add(pid as u8);
}
pkt
}
fn null_packet() -> [u8; 188] {
let mut pkt = [0u8; 188];
pkt[0] = 0x47;
pkt[1] = 0x1F;
pkt[2] = 0xFF;
pkt[3] = 0x10; for b in &mut pkt[4..] {
*b = 0xFF;
}
pkt
}
fn serialize_pat(pat: &PatSection) -> Vec<u8> {
let mut buf = vec![0u8; pat.serialized_len()];
pat.serialize_into(&mut buf).expect("PAT serialize");
buf
}
fn serialize_pmt(pmt: &PmtSection<'_>) -> Vec<u8> {
let mut buf = vec![0u8; pmt.serialized_len()];
pmt.serialize_into(&mut buf).expect("PMT serialize");
buf
}
fn build_two_program_ts(cycles: usize) -> Vec<u8> {
let pat = PatSection {
transport_stream_id: 1,
version_number: 0,
current_next_indicator: true,
section_number: 0,
last_section_number: 0,
entries: vec![
PatEntry {
program_number: 1,
pid: PMT1_PID,
},
PatEntry {
program_number: 2,
pid: PMT2_PID,
},
],
};
let pat_section_bytes = serialize_pat(&pat);
let mut pat_pktz = SectionPacketiser::new(PAT_PID);
let pat_pkts = pat_pktz.packetise(&[&pat_section_bytes]);
let pmt1 = PmtSection::new(
1,
0,
true,
0,
0,
P1_PCR_PID,
DescriptorLoop::new(&[]),
vec![
PmtStream {
stream_type: StreamType::Mpeg2Video,
elementary_pid: P1_VIDEO_PID,
es_info: DescriptorLoop::new(&[]),
},
PmtStream {
stream_type: StreamType::Mpeg2Audio,
elementary_pid: P1_AUDIO_PID,
es_info: DescriptorLoop::new(&[]),
},
],
);
let pmt1_section_bytes = serialize_pmt(&pmt1);
let mut pmt1_pktz = SectionPacketiser::new(PMT1_PID);
let pmt1_pkts = pmt1_pktz.packetise(&[&pmt1_section_bytes]);
let pmt2 = PmtSection::new(
2,
0,
true,
0,
0,
P2_PCR_PID,
DescriptorLoop::new(&[]),
vec![
PmtStream {
stream_type: StreamType::Mpeg2Video,
elementary_pid: P2_VIDEO_PID,
es_info: DescriptorLoop::new(&[]),
},
PmtStream {
stream_type: StreamType::Mpeg2Audio,
elementary_pid: P2_AUDIO_PID,
es_info: DescriptorLoop::new(&[]),
},
],
);
let pmt2_section_bytes = serialize_pmt(&pmt2);
let mut pmt2_pktz = SectionPacketiser::new(PMT2_PID);
let pmt2_pkts = pmt2_pktz.packetise(&[&pmt2_section_bytes]);
let mut stream: Vec<u8> = Vec::new();
for pkt in &pat_pkts {
stream.extend_from_slice(pkt);
}
for pkt in &pmt1_pkts {
stream.extend_from_slice(pkt);
}
for pkt in &pmt2_pkts {
stream.extend_from_slice(pkt);
}
for i in 0..cycles {
let cc = (i as u8) & 0x0F;
stream.extend_from_slice(&dummy_es_packet(P1_VIDEO_PID, cc));
stream.extend_from_slice(&dummy_es_packet(P2_VIDEO_PID, cc));
stream.extend_from_slice(&dummy_es_packet(P1_AUDIO_PID, cc));
stream.extend_from_slice(&dummy_es_packet(P2_AUDIO_PID, cc));
stream.extend_from_slice(&null_packet());
}
stream
}
fn pid_from_packet(pkt: &[u8]) -> u16 {
(((pkt[1] & 0x1F) as u16) << 8) | pkt[2] as u16
}
fn all_pids_in(ts: &[u8]) -> std::collections::BTreeSet<u16> {
ts.chunks_exact(188).map(pid_from_packet).collect()
}
fn pid_count(ts: &[u8], pid: u16) -> usize {
ts.chunks_exact(188)
.filter(|pkt| pid_from_packet(pkt) == pid)
.count()
}
fn run(input: &[u8], cfg: PidFilter) -> Vec<u8> {
let mut engine = TsFix::builder()
.filter_pids(cfg)
.build()
.expect("build should not fail");
let mut output = Vec::with_capacity(input.len());
for chunk in input.chunks_exact(188) {
engine
.push(chunk, |pkt| output.extend_from_slice(pkt))
.expect("valid 188-byte packet");
}
engine.finish(|pkt| output.extend_from_slice(pkt));
output
}
#[test]
fn service_extract_program_1() {
let ts = build_two_program_ts(8);
let output = run(&ts, PidFilter::service(1));
assert_eq!(
output.len() % 188,
0,
"output must be aligned to 188-byte packets"
);
assert!(!output.is_empty(), "output must not be empty");
let pids = all_pids_in(&output);
assert!(pids.contains(&PAT_PID), "PAT (0x0000) must be present");
assert!(pids.contains(&PMT1_PID), "PMT1 (0x0100) must be present");
assert!(
pids.contains(&P1_VIDEO_PID),
"P1 video (0x0101) must be present"
);
assert!(
pids.contains(&P1_AUDIO_PID),
"P1 audio (0x0102) must be present"
);
assert!(
!pids.contains(&PMT2_PID),
"PMT2 (0x0200) must NOT be present"
);
assert!(
!pids.contains(&P2_VIDEO_PID),
"P2 video (0x0201) must NOT be present"
);
assert!(
!pids.contains(&P2_AUDIO_PID),
"P2 audio (0x0202) must NOT be present"
);
assert!(
!pids.contains(&NULL_PID),
"null (0x1FFF) must NOT be present"
);
let expected: std::collections::BTreeSet<u16> =
[PAT_PID, PMT1_PID, P1_VIDEO_PID, P1_AUDIO_PID].into();
assert_eq!(pids, expected, "output must contain exactly the P1 PIDs");
}
#[test]
fn service_extract_program_2() {
let ts = build_two_program_ts(8);
let output = run(&ts, PidFilter::service(2));
assert_eq!(output.len() % 188, 0);
assert!(!output.is_empty());
let pids = all_pids_in(&output);
assert!(pids.contains(&PAT_PID));
assert!(pids.contains(&PMT2_PID));
assert!(pids.contains(&P2_VIDEO_PID));
assert!(pids.contains(&P2_AUDIO_PID));
assert!(!pids.contains(&PMT1_PID));
assert!(!pids.contains(&P1_VIDEO_PID));
assert!(!pids.contains(&P1_AUDIO_PID));
assert!(!pids.contains(&NULL_PID));
let expected: std::collections::BTreeSet<u16> =
[PAT_PID, PMT2_PID, P2_VIDEO_PID, P2_AUDIO_PID].into();
assert_eq!(pids, expected);
}
#[test]
fn keep_set_single_pid() {
let ts = build_two_program_ts(8);
let output = run(&ts, PidFilter::keep([P1_VIDEO_PID]));
assert_eq!(output.len() % 188, 0);
assert!(!output.is_empty());
let pids = all_pids_in(&output);
assert!(pids.contains(&PAT_PID), "PAT must always be present");
assert!(
pids.contains(&P1_VIDEO_PID),
"kept PID 0x0101 must be present"
);
assert!(!pids.contains(&PMT1_PID));
assert!(!pids.contains(&PMT2_PID));
assert!(!pids.contains(&P1_AUDIO_PID));
assert!(!pids.contains(&P2_VIDEO_PID));
assert!(!pids.contains(&P2_AUDIO_PID));
assert!(!pids.contains(&NULL_PID));
let expected: std::collections::BTreeSet<u16> = [PAT_PID, P1_VIDEO_PID].into();
assert_eq!(pids, expected);
}
#[test]
fn keep_set_empty_keeps_pat() {
let ts = build_two_program_ts(4);
let output = run(&ts, PidFilter::keep([]));
let pids = all_pids_in(&output);
let expected: std::collections::BTreeSet<u16> = [PAT_PID].into();
assert_eq!(pids, expected, "only PAT should survive an empty keep-set");
}
#[test]
fn interleaved_programs_are_truly_filtered() {
let ts = build_two_program_ts(4);
let packets: Vec<u16> = ts.chunks_exact(188).map(pid_from_packet).collect();
let first_p2_video = packets.iter().position(|&p| p == P2_VIDEO_PID);
let last_p1_audio = packets.iter().rposition(|&p| p == P1_AUDIO_PID);
assert!(
first_p2_video.is_some() && last_p1_audio.is_some(),
"fixture must contain both programs' ES packets"
);
assert!(
first_p2_video.unwrap() < last_p1_audio.unwrap(),
"packets must be interleaved (P2 video before last P1 audio)"
);
let output = run(&ts, PidFilter::service(1));
assert_eq!(pid_count(&output, P2_VIDEO_PID), 0);
assert_eq!(pid_count(&output, P2_AUDIO_PID), 0);
assert_eq!(pid_count(&output, PMT2_PID), 0);
assert!(pid_count(&output, P1_VIDEO_PID) > 0);
assert!(pid_count(&output, P1_AUDIO_PID) > 0);
}
#[test]
fn noop_fails_pid_exclusivity() {
let ts = build_two_program_ts(4);
let mut engine = TsFix::builder().build().expect("identity build");
let mut output = Vec::with_capacity(ts.len());
for chunk in ts.chunks_exact(188) {
engine
.push(chunk, |pkt| output.extend_from_slice(pkt))
.unwrap();
}
engine.finish(|pkt| output.extend_from_slice(pkt));
let pids = all_pids_in(&output);
assert!(
pids.contains(&PMT2_PID) && pids.contains(&P2_VIDEO_PID),
"identity engine must keep all PIDs — proves filter test would catch a broken impl"
);
}
#[test]
fn output_pmt_still_parses() {
let ts = build_two_program_ts(4);
let output = run(&ts, PidFilter::service(1));
let pmt1_pkt = output
.chunks_exact(188)
.find(|pkt| pid_from_packet(pkt) == PMT1_PID)
.expect("PMT1 packet must be in output");
let pusi = (pmt1_pkt[1] & 0x40) != 0;
assert!(pusi, "first PMT1 packet should have PUSI set");
let pointer = pmt1_pkt[4] as usize;
let section_start = 5 + pointer;
let section_bytes = &pmt1_pkt[section_start..];
let pmt = PmtSection::parse(section_bytes).expect("output PMT must parse successfully");
assert_eq!(pmt.program_number, 1);
assert_eq!(pmt.pcr_pid, P1_PCR_PID);
assert_eq!(pmt.streams.len(), 2);
}