use std::io;
use std::net::{IpAddr, SocketAddr};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use log::{debug, info, warn};
use plist::{Dictionary, Value};
use tokio::io::AsyncReadExt;
use tokio::net::{TcpListener, UdpSocket};
use tokio::task::JoinHandle;
use crate::buffered::{packet_seq, split_blocks, AudioDecryptor};
use crate::decode::AacDecoder;
use crate::dmap;
use crate::events::{Event, EventSender};
use crate::player::{Player, PlayerSender};
use crate::sink::{AudioSink, SinkFactory};
const TYPE_REALTIME: u64 = 96;
const TYPE_BUFFERED: u64 = 103;
const DMAP_CONTENT_TYPE: &str = "application/x-dmap-tagged";
pub struct Session {
local_ip: IpAddr,
tasks: Vec<JoinHandle<()>>,
stream_key: Option<Vec<u8>>,
audio_format: Option<u64>,
stream_type: Option<u64>,
volume: f32,
sink_factory: SinkFactory,
events: EventSender,
session_active: bool,
pending_metadata: Option<Event>,
pending_artwork: Option<Event>,
player: Option<Player>,
player_control: Option<PlayerSender>,
flush_until_seq: Arc<AtomicU64>,
}
impl Session {
pub fn new(local_ip: IpAddr, sink_factory: SinkFactory, events: EventSender) -> Session {
Session {
local_ip,
tasks: Vec::new(),
stream_key: None,
audio_format: None,
stream_type: None,
volume: 0.0,
sink_factory,
events,
session_active: false,
pending_metadata: None,
pending_artwork: None,
player: None,
player_control: None,
flush_until_seq: Arc::new(AtomicU64::new(0)),
}
}
fn send_event(&self, event: Event) {
let _ = self.events.send(event);
}
pub async fn handle_setup(&mut self, body: &[u8]) -> io::Result<Vec<u8>> {
let request = Value::from_reader(io::Cursor::new(body))
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, format!("SETUP plist: {e}")))?;
let dict = request
.as_dictionary()
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "SETUP body not a dict"))?;
if let Some(streams) = dict.get("streams").and_then(|v| v.as_array()) {
self.setup_streams(streams).await
} else {
self.setup_timing(dict).await
}
}
async fn setup_timing(&mut self, dict: &Dictionary) -> io::Result<Vec<u8>> {
let timing = dict
.get("timingProtocol")
.and_then(|v| v.as_string())
.unwrap_or("(none)");
info!("SETUP phase 1: timingProtocol={timing}");
let listener = TcpListener::bind(SocketAddr::new(self.local_ip, 0)).await?;
let event_port = listener.local_addr()?.port();
info!("SETUP phase 1: event port {event_port}");
self.tasks.push(tokio::spawn(event_channel(listener)));
let self_ip = self.local_ip.to_string();
let mut peer_info = Dictionary::new();
peer_info.insert(
"Addresses".into(),
Value::Array(vec![Value::String(self_ip.clone())]),
);
peer_info.insert("ID".into(), Value::String(self_ip));
let mut response = Dictionary::new();
response.insert(
"eventPort".into(),
Value::Integer(u64::from(event_port).into()),
);
response.insert("timingPort".into(), Value::Integer(0u64.into()));
response.insert("timingPeerInfo".into(), Value::Dictionary(peer_info));
encode_plist(&response)
}
async fn setup_streams(&mut self, streams: &[Value]) -> io::Result<Vec<u8>> {
let stream = streams
.first()
.and_then(|v| v.as_dictionary())
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "empty streams array"))?;
let stream_type = stream.get("type").and_then(|v| v.as_unsigned_integer());
self.stream_type = stream_type;
self.audio_format = stream
.get("audioFormat")
.and_then(|v| v.as_unsigned_integer());
self.stream_key = stream
.get("shk")
.and_then(|v| v.as_data())
.map(<[u8]>::to_vec);
let spf = stream.get("spf").and_then(|v| v.as_unsigned_integer());
info!(
"SETUP phase 2: type={stream_type:?} audioFormat={:?} spf={spf:?} shk={}",
self.audio_format,
self.stream_key.as_ref().map_or(0, Vec::len)
);
let control = UdpSocket::bind(SocketAddr::new(self.local_ip, 0)).await?;
let control_port = control.local_addr()?.port();
self.tasks
.push(tokio::spawn(audio_channel(control, "control")));
let data_port = if stream_type == Some(TYPE_BUFFERED) {
self.start_buffered_audio().await?
} else {
let data = UdpSocket::bind(SocketAddr::new(self.local_ip, 0)).await?;
let port = data.local_addr()?.port();
self.tasks.push(tokio::spawn(audio_channel(data, "audio")));
port
};
info!("SETUP phase 2: data port {data_port}, control port {control_port}");
let mut stream_response = Dictionary::new();
stream_response.insert(
"type".into(),
Value::Integer(stream_type.unwrap_or(TYPE_REALTIME).into()),
);
stream_response.insert(
"dataPort".into(),
Value::Integer(u64::from(data_port).into()),
);
stream_response.insert(
"controlPort".into(),
Value::Integer(u64::from(control_port).into()),
);
if stream_type == Some(TYPE_BUFFERED) {
stream_response.insert(
"audioBufferSize".into(),
Value::Integer(8_388_608u64.into()),
);
}
let mut response = Dictionary::new();
response.insert(
"streams".into(),
Value::Array(vec![Value::Dictionary(stream_response)]),
);
encode_plist(&response)
}
async fn start_buffered_audio(&mut self) -> io::Result<u16> {
let listener = TcpListener::bind(SocketAddr::new(self.local_ip, 0)).await?;
let port = listener.local_addr()?.port();
let decryptor = self.stream_key.as_deref().and_then(AudioDecryptor::new);
let (rate, channels) = aac_params(self.audio_format);
self.flush_until_seq.store(0, Ordering::Relaxed);
match (decryptor, AacDecoder::new(rate, channels)) {
(Some(decryptor), Ok(decoder)) => {
self.session_active = true;
self.send_event(Event::SessionStarted { rate, channels });
let latched = [self.pending_metadata.take(), self.pending_artwork.take()];
for event in latched.into_iter().flatten() {
self.send_event(event);
}
let sink: Box<dyn AudioSink> = (self.sink_factory)(rate, channels);
let player = Player::spawn(sink);
let sender = player.sender();
self.player_control = Some(player.sender());
self.player = Some(player);
let max_queued = crate::player::max_queued_samples(rate, channels);
self.tasks.push(tokio::spawn(buffered_audio(
listener,
decryptor,
decoder,
sender,
max_queued,
self.flush_until_seq.clone(),
)));
info!("buffered audio: TCP data port {port}, {rate} Hz {channels}ch");
}
(None, _) => {
warn!("buffered audio: missing/invalid shk; draining without decode");
self.tasks.push(tokio::spawn(drain_tcp(listener)));
}
(_, Err(e)) => {
warn!("buffered audio: decoder init failed ({e}); draining");
self.tasks.push(tokio::spawn(drain_tcp(listener)));
}
}
Ok(port)
}
pub fn set_rate_anchor(&mut self, body: &[u8]) {
let Some((rate, rtp)) = parse_rate_anchor(body) else {
warn!("SETRATEANCHORTIME: could not parse body");
return;
};
debug!("SETRATEANCHORTIME rate={rate} rtpTime={rtp}");
if let Some(ctrl) = &self.player_control {
ctrl.set_paused(rate == 0);
}
self.send_event(Event::Paused(rate == 0));
}
pub fn flush(&mut self, body: &[u8]) {
let boundary = parse_flush_until_seq(body);
match boundary {
Some(seq) => {
self.flush_until_seq.store(seq, Ordering::Relaxed);
debug!("FLUSHBUFFERED until seq {seq}");
}
None => debug!("FLUSHBUFFERED (no seq boundary)"),
}
if let Some(ctrl) = &self.player_control {
ctrl.flush(boundary);
}
self.send_event(Event::Flushed);
}
pub fn get_parameter(&self, body: &[u8]) -> Vec<u8> {
let query = String::from_utf8_lossy(body);
if query.trim() == "volume" {
format!("volume: {:.6}\r\n", self.volume).into_bytes()
} else {
debug!("GET_PARAMETER for unknown parameter: {query:?}");
Vec::new()
}
}
pub fn set_parameter(&mut self, content_type: Option<&str>, body: &[u8]) {
let media_type = content_type.map(|ct| ct.split(';').next().unwrap_or(ct).trim());
match media_type {
Some(ct) if ct.eq_ignore_ascii_case(DMAP_CONTENT_TYPE) => self.set_metadata(body),
Some(ct)
if ct
.get(..6)
.is_some_and(|p| p.eq_ignore_ascii_case("image/")) =>
{
self.set_artwork(ct, body)
}
_ => self.set_text_parameters(body),
}
}
fn set_text_parameters(&mut self, body: &[u8]) {
let text = String::from_utf8_lossy(body);
for line in text.lines() {
if let Some(v) = line.trim().strip_prefix("volume:") {
if let Ok(db) = v.trim().parse::<f32>() {
self.volume = db;
debug!("SET_PARAMETER volume {db} dB");
self.send_event(Event::Volume { db });
}
}
}
}
fn set_metadata(&mut self, body: &[u8]) {
let Some(meta) = dmap::parse(body) else {
debug!(
"SET_PARAMETER metadata: unrecognized DMAP payload ({} bytes)",
body.len()
);
return;
};
debug!(
"SET_PARAMETER metadata: title={:?} artist={:?} album={:?}",
meta.title, meta.artist, meta.album
);
self.send_session_event(Event::Metadata {
title: meta.title,
artist: meta.artist,
album: meta.album,
});
}
fn set_artwork(&mut self, content_type: &str, body: &[u8]) {
debug!(
"SET_PARAMETER artwork: {content_type}, {} bytes",
body.len()
);
self.send_session_event(Event::Artwork {
content_type: content_type.to_string(),
data: body.to_vec(),
});
}
fn send_session_event(&mut self, event: Event) {
if self.session_active {
self.send_event(event);
} else if matches!(event, Event::Artwork { .. }) {
self.pending_artwork = Some(event);
} else {
self.pending_metadata = Some(event);
}
}
pub fn ack(&self, method: &str) {
debug!("ack {method}");
}
pub fn teardown(&mut self) {
debug!("ack TEARDOWN");
self.end_session();
}
fn end_session(&mut self) {
if self.session_active {
self.session_active = false;
self.send_event(Event::SessionEnded);
}
}
}
impl Drop for Session {
fn drop(&mut self) {
self.end_session();
for task in self.tasks.drain(..) {
task.abort();
}
}
}
fn encode_plist(dict: &Dictionary) -> io::Result<Vec<u8>> {
let mut buf = Vec::new();
Value::Dictionary(dict.clone())
.to_writer_binary(&mut buf)
.map_err(|e| io::Error::other(format!("plist encode: {e}")))?;
Ok(buf)
}
async fn event_channel(listener: TcpListener) {
let Ok((mut stream, peer)) = listener.accept().await else {
return;
};
debug!("event channel connected from {peer}");
let mut buf = [0u8; 4096];
use tokio::io::AsyncReadExt;
while let Ok(n) = stream.read(&mut buf).await {
if n == 0 {
break;
}
debug!("event: {n} bytes");
}
}
async fn audio_channel(socket: UdpSocket, label: &'static str) {
let mut buf = vec![0u8; 16 * 1024];
let mut count: u64 = 0;
loop {
match socket.recv(&mut buf).await {
Ok(n) => {
count += 1;
if count <= 3 || count.is_multiple_of(250) {
info!("{label}: {count} packets, last {n} bytes");
}
}
Err(e) => {
warn!("{label} socket error: {e}");
return;
}
}
}
}
fn aac_params(_audio_format: Option<u64>) -> (u32, u8) {
(44100, 2)
}
fn parse_rate_anchor(body: &[u8]) -> Option<(u64, u64)> {
let value = Value::from_reader(io::Cursor::new(body)).ok()?;
let dict = value.as_dictionary()?;
let rate = dict.get("rate").and_then(int_field)?;
let rtp = dict.get("rtpTime").and_then(int_field).unwrap_or(0);
Some((rate, rtp))
}
fn int_field(v: &Value) -> Option<u64> {
v.as_unsigned_integer()
.or_else(|| v.as_signed_integer().map(|s| s as u64))
}
fn parse_flush_until_seq(body: &[u8]) -> Option<u64> {
let value = Value::from_reader(io::Cursor::new(body)).ok()?;
value
.as_dictionary()?
.get("flushUntilSeq")
.and_then(int_field)
}
async fn buffered_audio(
listener: TcpListener,
decryptor: AudioDecryptor,
mut decoder: AacDecoder,
player: PlayerSender,
max_queued: usize,
flush_until_seq: Arc<AtomicU64>,
) {
let Ok((mut stream, peer)) = listener.accept().await else {
return;
};
info!("buffered audio connected from {peer}");
let mut buf: Vec<u8> = Vec::new();
let mut chunk = vec![0u8; 64 * 1024];
let mut decrypt_failures: u64 = 0;
let mut skipped: u64 = 0;
loop {
while player.pending_samples() > max_queued {
tokio::time::sleep(Duration::from_millis(5)).await;
}
match stream.read(&mut chunk).await {
Ok(0) => break,
Ok(n) => {
buf.extend_from_slice(&chunk[..n]);
let (packets, used) = split_blocks(&buf);
let owned: Vec<Vec<u8>> = packets.iter().map(|p| p.to_vec()).collect();
buf.drain(..used);
for packet in owned {
if let Some(seq) = packet_seq(&packet) {
if skip_before_boundary(&flush_until_seq, seq) {
skipped += 1;
if skipped <= 3 || skipped.is_multiple_of(2000) {
debug!("buffered audio: skipping seq {seq}");
}
continue;
}
}
let Some(audio) = decryptor.decrypt(&packet) else {
decrypt_failures += 1;
if decrypt_failures <= 3 {
debug!("buffered audio: packet decrypt failed");
}
continue;
};
match decoder.decode(&audio.payload) {
Ok(pcm) if !pcm.is_empty() => player.play(u64::from(audio.seq), pcm),
Ok(_) => {}
Err(e) => debug!("buffered audio: decode error: {e}"),
}
}
}
Err(e) => {
warn!("buffered audio read error: {e}");
break;
}
}
}
info!("buffered audio disconnected ({decrypt_failures} decrypt failures, {skipped} skipped)");
}
fn skip_before_boundary(flush_until_seq: &AtomicU64, seq: u32) -> bool {
let boundary = flush_until_seq.load(Ordering::Relaxed);
if boundary == 0 {
return false;
}
if u64::from(seq) < boundary {
return true;
}
let _ = flush_until_seq.compare_exchange(boundary, 0, Ordering::Relaxed, Ordering::Relaxed);
false
}
async fn drain_tcp(listener: TcpListener) {
let Ok((mut stream, _)) = listener.accept().await else {
return;
};
let mut buf = vec![0u8; 64 * 1024];
while let Ok(n) = stream.read(&mut buf).await {
if n == 0 {
break;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::Ipv4Addr;
use tokio::sync::mpsc::UnboundedReceiver;
fn local() -> IpAddr {
IpAddr::V4(Ipv4Addr::LOCALHOST)
}
struct TestSink;
impl AudioSink for TestSink {
fn write(&mut self, _pcm: &[i16]) {}
fn flush(&mut self) {}
}
fn session() -> (Session, UnboundedReceiver<Event>) {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
let factory: SinkFactory = Arc::new(|_, _| Box::new(TestSink));
(Session::new(local(), factory, tx), rx)
}
#[tokio::test]
async fn phase1_response_has_event_and_timing_ports() {
let (mut session, _events) = session();
let mut dict = Dictionary::new();
dict.insert("timingProtocol".into(), Value::String("PTP".into()));
let body = encode_plist(&dict).unwrap();
let response = session.handle_setup(&body).await.unwrap();
let value = Value::from_reader(io::Cursor::new(response)).unwrap();
let d = value.as_dictionary().unwrap();
assert!(d.get("eventPort").unwrap().as_unsigned_integer().unwrap() > 0);
assert_eq!(d.get("timingPort").unwrap().as_unsigned_integer(), Some(0));
assert!(d.get("timingPeerInfo").unwrap().as_dictionary().is_some());
}
#[tokio::test]
async fn phase2_response_binds_ports_and_echoes_type() {
let (mut session, mut events) = session();
let mut stream = Dictionary::new();
stream.insert("type".into(), Value::Integer(TYPE_BUFFERED.into()));
stream.insert("audioFormat".into(), Value::Integer(0x40000u64.into()));
stream.insert("shk".into(), Value::Data(vec![7u8; 32]));
let mut dict = Dictionary::new();
dict.insert(
"streams".into(),
Value::Array(vec![Value::Dictionary(stream)]),
);
let body = encode_plist(&dict).unwrap();
let response = session.handle_setup(&body).await.unwrap();
let value = Value::from_reader(io::Cursor::new(response)).unwrap();
let streams = value
.as_dictionary()
.unwrap()
.get("streams")
.unwrap()
.as_array()
.unwrap();
let s = streams[0].as_dictionary().unwrap();
assert_eq!(
s.get("type").unwrap().as_unsigned_integer(),
Some(TYPE_BUFFERED)
);
assert!(s.get("dataPort").unwrap().as_unsigned_integer().unwrap() > 0);
assert!(s.get("controlPort").unwrap().as_unsigned_integer().unwrap() > 0);
assert!(s.get("audioBufferSize").is_some());
assert_eq!(session.stream_key, Some(vec![7u8; 32]));
assert_eq!(
events.try_recv(),
Ok(Event::SessionStarted {
rate: 44100,
channels: 2
})
);
drop(session);
assert_eq!(events.try_recv(), Ok(Event::SessionEnded));
assert!(events.try_recv().is_err());
}
#[tokio::test]
async fn phase_detection_uses_streams_presence() {
let (mut s1, _events) = session();
let empty = encode_plist(&Dictionary::new()).unwrap();
let r1 = s1.handle_setup(&empty).await.unwrap();
assert!(Value::from_reader(io::Cursor::new(r1))
.unwrap()
.as_dictionary()
.unwrap()
.contains_key("eventPort"));
}
#[test]
fn volume_query_returns_current_volume() {
let (mut session, mut events) = session();
assert_eq!(
session.get_parameter(b"volume\r\n"),
b"volume: 0.000000\r\n"
);
session.set_parameter(Some("text/parameters"), b"volume: -12.5\r\n");
assert_eq!(
session.get_parameter(b"volume\r\n"),
b"volume: -12.500000\r\n"
);
assert_eq!(events.try_recv(), Ok(Event::Volume { db: -12.5 }));
assert!(session.get_parameter(b"progress\r\n").is_empty());
}
fn dmap_entry(tag: &[u8; 4], payload: &[u8]) -> Vec<u8> {
let mut e = tag.to_vec();
e.extend_from_slice(&(payload.len() as u32).to_be_bytes());
e.extend_from_slice(payload);
e
}
fn dmap_track(title: &str) -> Vec<u8> {
let children = [
dmap_entry(b"minm", title.as_bytes()),
dmap_entry(b"asar", b"Artist"),
dmap_entry(b"asal", b"Album"),
]
.concat();
dmap_entry(b"mlit", &children)
}
async fn start_stream(session: &mut Session) {
let mut stream = Dictionary::new();
stream.insert("type".into(), Value::Integer(TYPE_BUFFERED.into()));
stream.insert("shk".into(), Value::Data(vec![7u8; 32]));
let mut dict = Dictionary::new();
dict.insert(
"streams".into(),
Value::Array(vec![Value::Dictionary(stream)]),
);
let body = encode_plist(&dict).unwrap();
session.handle_setup(&body).await.unwrap();
}
#[tokio::test]
async fn metadata_and_artwork_reach_the_host_mid_session() {
let (mut session, mut events) = session();
start_stream(&mut session).await;
assert!(matches!(
events.try_recv(),
Ok(Event::SessionStarted { .. })
));
session.set_parameter(Some(DMAP_CONTENT_TYPE), &dmap_track("Song"));
assert_eq!(
events.try_recv(),
Ok(Event::Metadata {
title: Some("Song".into()),
artist: Some("Artist".into()),
album: Some("Album".into()),
})
);
session.set_parameter(Some("image/png"), b"\x89PNG");
assert_eq!(
events.try_recv(),
Ok(Event::Artwork {
content_type: "image/png".into(),
data: b"\x89PNG".to_vec(),
})
);
session.set_parameter(Some("image/none"), b"");
assert_eq!(
events.try_recv(),
Ok(Event::Artwork {
content_type: "image/none".into(),
data: Vec::new(),
})
);
}
#[tokio::test]
async fn early_metadata_is_latched_until_session_start() {
let (mut session, mut events) = session();
session.set_parameter(Some(DMAP_CONTENT_TYPE), &dmap_track("First"));
session.set_parameter(Some(DMAP_CONTENT_TYPE), &dmap_track("Second"));
session.set_parameter(Some("image/jpeg"), b"JPEG");
assert!(events.try_recv().is_err());
start_stream(&mut session).await;
assert!(matches!(
events.try_recv(),
Ok(Event::SessionStarted { .. })
));
assert!(matches!(
events.try_recv(),
Ok(Event::Metadata { title: Some(t), .. }) if t == "Second"
));
assert!(matches!(events.try_recv(), Ok(Event::Artwork { .. })));
assert!(events.try_recv().is_err());
}
#[tokio::test]
async fn malformed_metadata_never_reaches_the_host() {
let (mut session, mut events) = session();
start_stream(&mut session).await;
assert!(matches!(
events.try_recv(),
Ok(Event::SessionStarted { .. })
));
session.set_parameter(Some(DMAP_CONTENT_TYPE), b"");
session.set_parameter(Some(DMAP_CONTENT_TYPE), b"garbage, not dmap");
session.set_parameter(Some(DMAP_CONTENT_TYPE), b"mlit\x00\x00\xff\xff");
assert!(events.try_recv().is_err());
session.set_parameter(Some("text/parameters"), b"volume: -6.0\r\n");
assert_eq!(events.try_recv(), Ok(Event::Volume { db: -6.0 }));
}
#[test]
fn parses_real_setrateanchortime() {
let mut dict = Dictionary::new();
dict.insert("rate".into(), Value::Integer(1u64.into()));
dict.insert(
"networkTimeTimelineID".into(),
Value::Integer((-2116301217048756216i64).into()),
);
dict.insert("networkTimeSecs".into(), Value::Integer(1323152u64.into()));
dict.insert(
"networkTimeFrac".into(),
Value::Integer(6275326383463858176u64.into()),
);
dict.insert("networkTimeFlags".into(), Value::Integer(0u64.into()));
dict.insert("rtpTime".into(), Value::Integer(3174381381u64.into()));
let body = encode_plist(&dict).unwrap();
assert_eq!(parse_rate_anchor(&body), Some((1, 3174381381)));
}
#[test]
fn flush_boundary_is_not_sticky() {
let boundary = AtomicU64::new(100);
assert!(skip_before_boundary(&boundary, 50));
assert!(skip_before_boundary(&boundary, 99));
assert!(!skip_before_boundary(&boundary, 100));
assert_eq!(boundary.load(Ordering::Relaxed), 0);
assert!(!skip_before_boundary(&boundary, 50));
assert!(!skip_before_boundary(&AtomicU64::new(0), 7));
}
#[test]
fn parses_real_flushbuffered() {
let mut dict = Dictionary::new();
dict.insert("flushUntilSeq".into(), Value::Integer(5179978u64.into()));
dict.insert("flushUntilTS".into(), Value::Integer(2204469244u64.into()));
let body = encode_plist(&dict).unwrap();
assert_eq!(parse_flush_until_seq(&body), Some(5179978));
let empty = encode_plist(&Dictionary::new()).unwrap();
assert_eq!(parse_flush_until_seq(&empty), None);
}
#[test]
fn rate_anchor_pause_and_missing_fields() {
let mut dict = Dictionary::new();
dict.insert("rate".into(), Value::Integer(0u64.into()));
let body = encode_plist(&dict).unwrap();
assert_eq!(parse_rate_anchor(&body), Some((0, 0)));
let empty = encode_plist(&Dictionary::new()).unwrap();
assert_eq!(parse_rate_anchor(&empty), None);
}
}