use super::meta;
use crate::disc::DiscTitle;
use std::io::{self, BufReader, BufWriter, Write};
use std::net::{IpAddr, TcpListener, TcpStream, ToSocketAddrs};
const NET_BUF_SIZE: usize = 256 * 1024;
pub(crate) fn is_blocked_ip(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => {
v4.is_loopback()
|| v4.is_private()
|| v4.is_link_local()
|| v4.is_unspecified()
|| v4.is_multicast()
|| v4.is_broadcast()
}
IpAddr::V6(v6) => {
v6.is_loopback()
|| v6.is_unspecified()
|| v6.is_multicast()
|| (v6.segments()[0] & 0xfe00) == 0xfc00
|| (v6.segments()[0] & 0xffc0) == 0xfe80
}
}
}
fn resolve_allowed_addr(addr: &str) -> io::Result<std::net::SocketAddr> {
addr.to_socket_addrs()?
.find(|sa| !is_blocked_ip(sa.ip()))
.ok_or_else(|| {
crate::error::Error::NetworkAddrBlocked {
addr: addr.to_string(),
}
.into()
})
}
enum Mode {
Write {
writer: BufWriter<TcpStream>,
header_written: bool,
},
Read {
reader: BufReader<TcpStream>,
},
}
pub struct NetworkStream {
disc_title: DiscTitle,
mode: Mode,
}
impl NetworkStream {
pub fn connect(addr: &str) -> io::Result<Self> {
Self::connect_vetted(addr, true)
}
fn connect_vetted(addr: &str, vet: bool) -> io::Result<Self> {
let stream = if vet {
let vetted = resolve_allowed_addr(addr)?;
TcpStream::connect(vetted)?
} else {
TcpStream::connect(addr)?
};
stream.set_nodelay(true)?;
Ok(Self {
disc_title: DiscTitle::empty(),
mode: Mode::Write {
writer: BufWriter::with_capacity(NET_BUF_SIZE, stream),
header_written: false,
},
})
}
pub fn meta(mut self, dt: &DiscTitle) -> Self {
self.disc_title = dt.clone();
self
}
pub fn listen(addr: &str) -> io::Result<Self> {
Self::accept_from(TcpListener::bind(addr)?)
}
pub fn accept_from(listener: TcpListener) -> io::Result<Self> {
let (stream, _peer) = listener.accept()?;
stream.set_nodelay(true)?;
let mut reader = BufReader::with_capacity(NET_BUF_SIZE, stream);
let disc_title = meta::read_header(&mut reader)?
.ok_or_else(|| -> io::Error { crate::error::Error::NoMetadata.into() })?
.to_title();
Ok(Self {
disc_title,
mode: Mode::Read { reader },
})
}
}
fn ensure_header_written(
writer: &mut BufWriter<TcpStream>,
header_written: &mut bool,
disc_title: &DiscTitle,
) -> io::Result<()> {
if !*header_written {
let m = meta::M2tsMeta::from_title(disc_title);
meta::write_header(writer, &m)?;
*header_written = true;
}
Ok(())
}
impl crate::pes::Stream for NetworkStream {
fn read(&mut self) -> io::Result<Option<crate::pes::PesFrame>> {
match &mut self.mode {
Mode::Read { reader } => crate::pes::PesFrame::deserialize(reader),
_ => Err(crate::error::Error::StreamWriteOnly.into()),
}
}
fn write(&mut self, frame: &crate::pes::PesFrame) -> io::Result<()> {
match &mut self.mode {
Mode::Write {
writer,
header_written,
} => {
ensure_header_written(writer, header_written, &self.disc_title)?;
frame.serialize(writer)
}
_ => Err(crate::error::Error::StreamReadOnly.into()),
}
}
fn finish(&mut self) -> io::Result<()> {
if let Mode::Write {
writer,
header_written,
} = &mut self.mode
{
ensure_header_written(writer, header_written, &self.disc_title)?;
writer.flush()?;
writer.get_ref().shutdown(std::net::Shutdown::Write)?;
}
Ok(())
}
fn info(&self) -> &DiscTitle {
&self.disc_title
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::disc::{
AudioChannels, AudioStream, Codec, ColorSpace, ContentFormat, FrameRate, HdrFormat,
Resolution, SampleRate, Stream, VideoStream,
};
use std::net::TcpListener;
#[test]
fn is_blocked_ip_rejects_internal_targets() {
use std::net::{Ipv4Addr, Ipv6Addr};
let v4 = |a, b, c, d| IpAddr::V4(Ipv4Addr::new(a, b, c, d));
let blocked: &[(IpAddr, &str)] = &[
(v4(127, 0, 0, 1), "loopback"),
(v4(127, 10, 20, 30), "loopback /8"),
(v4(10, 0, 0, 1), "private 10/8"),
(v4(172, 16, 5, 5), "private 172.16/12"),
(v4(192, 168, 1, 1), "private 192.168/16"),
(v4(169, 254, 10, 10), "link-local"),
(v4(0, 0, 0, 0), "unspecified"),
(v4(224, 0, 0, 1), "multicast"),
(v4(255, 255, 255, 255), "broadcast"),
(IpAddr::V6(Ipv6Addr::LOCALHOST), "loopback v6"),
(IpAddr::V6(Ipv6Addr::UNSPECIFIED), "unspecified v6"),
(
IpAddr::V6(Ipv6Addr::new(0xfc00, 0, 0, 0, 0, 0, 0, 1)),
"ULA",
),
(
IpAddr::V6(Ipv6Addr::new(0xfd12, 0x3456, 0, 0, 0, 0, 0, 1)),
"ULA",
),
(
IpAddr::V6(Ipv6Addr::new(0xfe80, 0, 0, 0, 0, 0, 0, 1)),
"link-local v6",
),
(
IpAddr::V6(Ipv6Addr::new(0xff02, 0, 0, 0, 0, 0, 0, 1)),
"multicast v6",
),
];
for (ip, label) in blocked {
assert!(is_blocked_ip(*ip), "{label} ({ip}) must be blocked");
}
let allowed: &[(IpAddr, &str)] = &[
(v4(8, 8, 8, 8), "public dns"),
(v4(1, 1, 1, 1), "public dns"),
(v4(93, 184, 216, 34), "example.com"),
(
IpAddr::V6(Ipv6Addr::new(0x2606, 0x2800, 0x220, 1, 0, 0, 0, 1)),
"public v6",
),
];
for (ip, label) in allowed {
assert!(!is_blocked_ip(*ip), "{label} ({ip}) must be allowed");
}
}
#[test]
fn connect_refuses_blocked_loopback_target() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let err = match NetworkStream::connect(&addr.to_string()) {
Ok(_) => panic!("loopback target must be refused by the SSRF guard"),
Err(e) => e,
};
assert_eq!(err.kind(), io::ErrorKind::PermissionDenied);
}
fn sample_title() -> DiscTitle {
DiscTitle {
playlist: "NetworkTest".into(),
playlist_id: 1,
duration_secs: 3600.0,
size_bytes: 0,
clips: Vec::new(),
streams: vec![
Stream::Video(VideoStream {
pid: 0x1011,
codec: Codec::Hevc,
resolution: Resolution::R2160p,
frame_rate: FrameRate::F23_976,
hdr: HdrFormat::Hdr10,
color_space: ColorSpace::Bt2020,
secondary: false,
label: "Main".into(),
}),
Stream::Audio(AudioStream {
pid: 0x1100,
codec: Codec::TrueHd,
channels: AudioChannels::Surround71,
language: "eng".into(),
sample_rate: SampleRate::S48,
secondary: false,
purpose: crate::disc::LabelPurpose::Normal,
label: "English".into(),
}),
],
chapters: Vec::new(),
extents: Vec::new(),
content_format: ContentFormat::BdTs,
codec_privates: Vec::new(),
}
}
#[test]
fn network_pes_roundtrip() {
use crate::pes;
use std::sync::mpsc;
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let (addr_tx, addr_rx) = mpsc::channel();
let handle = std::thread::spawn(move || {
addr_tx.send(addr).unwrap();
let mut ns = NetworkStream::accept_from(listener).unwrap();
let info = pes::Stream::info(&ns).clone();
let mut frames = Vec::new();
while let Ok(Some(f)) = pes::Stream::read(&mut ns) {
frames.push(f);
}
(info, frames)
});
let addr = addr_rx.recv().unwrap();
let dt = sample_title();
let mut writer = NetworkStream::connect_vetted(&addr.to_string(), false)
.unwrap()
.meta(&dt);
let frame = pes::PesFrame {
track: 0,
pts: 90000,
keyframe: true,
data: vec![0x47; 192],
duration_ns: None,
};
pes::Stream::write(&mut writer, &frame).unwrap();
pes::Stream::finish(&mut writer).unwrap();
let (info, frames) = handle.join().unwrap();
assert_eq!(info.playlist, "NetworkTest");
assert_eq!(info.streams.len(), 2);
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].track, 0);
assert_eq!(frames[0].pts, 90000);
}
#[test]
fn network_zero_frame_finish_still_sends_header() {
use crate::pes;
use std::sync::mpsc;
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let (addr_tx, addr_rx) = mpsc::channel();
let handle = std::thread::spawn(move || {
addr_tx.send(addr).unwrap();
let ns = NetworkStream::accept_from(listener).unwrap();
pes::Stream::info(&ns).playlist.clone()
});
let addr = addr_rx.recv().unwrap();
let dt = sample_title();
let mut writer = NetworkStream::connect_vetted(&addr.to_string(), false)
.unwrap()
.meta(&dt);
pes::Stream::finish(&mut writer).unwrap();
let playlist = handle.join().unwrap();
assert_eq!(
playlist, "NetworkTest",
"zero-frame finish() must still deliver the metadata header"
);
}
#[test]
fn network_empty_addr_errors() {
let result = NetworkStream::connect("");
assert!(result.is_err());
}
#[test]
fn network_no_port_errors() {
let result = NetworkStream::connect("127.0.0.1");
assert!(result.is_err());
}
fn spawn_reader() -> (
std::net::SocketAddr,
std::thread::JoinHandle<(DiscTitle, Vec<crate::pes::PesFrame>)>,
) {
use crate::pes;
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let handle = std::thread::spawn(move || {
let mut ns = NetworkStream::accept_from(listener).unwrap();
let info = pes::Stream::info(&ns).clone();
let mut frames = Vec::new();
while let Ok(Some(f)) = pes::Stream::read(&mut ns) {
frames.push(f);
}
(info, frames)
});
(addr, handle)
}
#[test]
fn write_on_read_side_is_read_only_error() {
use crate::pes;
let (addr, handle) = spawn_reader();
let dt = sample_title();
let mut writer = NetworkStream::connect_vetted(&addr.to_string(), false)
.unwrap()
.meta(&dt);
pes::Stream::finish(&mut writer).unwrap();
let (_info, _frames) = handle.join().unwrap();
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr2 = listener.local_addr().unwrap();
let h = std::thread::spawn(move || {
let mut ns = NetworkStream::accept_from(listener).unwrap();
let frame = pes::PesFrame {
track: 0,
pts: 0,
keyframe: true,
data: vec![0u8; 8],
duration_ns: None,
};
let err = pes::Stream::write(&mut ns, &frame).expect_err("read side write must error");
err.kind()
});
let mut w2 = NetworkStream::connect_vetted(&addr2.to_string(), false)
.unwrap()
.meta(&dt);
pes::Stream::finish(&mut w2).unwrap();
let kind = h.join().unwrap();
assert_eq!(kind, io::ErrorKind::Unsupported);
}
#[test]
fn read_on_write_side_is_write_only_error() {
use crate::pes;
let (addr, handle) = spawn_reader();
let dt = sample_title();
let mut writer = NetworkStream::connect_vetted(&addr.to_string(), false)
.unwrap()
.meta(&dt);
let err = pes::Stream::read(&mut writer).expect_err("write side read must error");
assert_eq!(err.kind(), io::ErrorKind::Unsupported);
pes::Stream::finish(&mut writer).unwrap();
let _ = handle.join().unwrap();
}
#[test]
fn header_written_once_then_all_frames_roundtrip() {
use crate::pes;
let (addr, handle) = spawn_reader();
let dt = sample_title();
let mut writer = NetworkStream::connect_vetted(&addr.to_string(), false)
.unwrap()
.meta(&dt);
for i in 0..5u8 {
let frame = pes::PesFrame {
track: (i % 2) as usize,
pts: i as i64 * 90_000,
keyframe: i == 0,
data: vec![i; 100 + i as usize],
duration_ns: None,
};
pes::Stream::write(&mut writer, &frame).unwrap();
}
pes::Stream::finish(&mut writer).unwrap();
let (info, frames) = handle.join().unwrap();
assert_eq!(info.streams.len(), 2);
assert_eq!(frames.len(), 5);
for (i, f) in frames.iter().enumerate() {
assert_eq!(f.pts, i as i64 * 90_000, "frame {i} pts");
assert_eq!(f.data.len(), 100 + i, "frame {i} payload length");
assert!(
f.data.iter().all(|&b| b == i as u8),
"frame {i} payload bytes"
);
}
}
#[test]
fn receiver_title_comes_from_sender_header() {
use crate::pes;
let (addr, handle) = spawn_reader();
let mut dt = sample_title();
dt.playlist = "SenderControlled".into();
dt.playlist_id = 42;
let mut writer = NetworkStream::connect_vetted(&addr.to_string(), false)
.unwrap()
.meta(&dt);
pes::Stream::finish(&mut writer).unwrap();
let (info, _frames) = handle.join().unwrap();
assert_eq!(info.playlist, "SenderControlled");
assert_eq!(
info.streams.len(),
2,
"stream descriptors round-trip via header"
);
}
#[test]
fn accept_from_rejects_stream_without_fmkv_header() {
use std::io::Write as _;
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let handle = std::thread::spawn(move || {
let mut s = TcpStream::connect(addr).unwrap();
s.write_all(&[0x47u8; 64]).unwrap(); s.shutdown(std::net::Shutdown::Both).unwrap();
});
let err = match NetworkStream::accept_from(listener) {
Ok(_) => panic!("missing FMKV header must error, not silently accept"),
Err(e) => e,
};
assert_eq!(err.kind(), io::ErrorKind::InvalidInput);
handle.join().unwrap();
}
}