#![cfg(all(not(target_arch = "wasm32"), feature = "transport-moq"))]
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU8, Ordering};
use std::sync::Arc;
use std::time::Duration;
use anyhow::{anyhow, bail, ensure, Context, Result};
use bytes::Bytes;
use moq_transport::coding::TrackNamespace;
use moq_transport::serve::{self, TrackReaderMode};
use tokio::sync::{Mutex, RwLock};
use super::{NativeMoQMessageHandler, NativeMoQState, NativeMoQTerminalHandler};
#[cfg(feature = "test-harness")]
pub mod test_relay;
const MOQ_SETUP_TIMEOUT: Duration = Duration::from_secs(10);
const MOQ_SUBSCRIBE_RETRY_DELAY: Duration = Duration::from_millis(250);
const MOQ_MAX_OBJECT_PAYLOAD_BYTES: usize = 4 * 1024 * 1024;
const MOQ_MAX_ACCESS_TOKEN_BYTES: usize = 16 * 1024;
const MOQ_BROADCAST_TRACK: &str = "data/broadcast";
type Draft14TrackWriters = HashMap<String, serve::SubgroupsWriter>;
#[derive(Clone)]
pub struct NativeMoQSession {
local_node_id: String,
remote_node_id: String,
relay_url: String,
access_token: Option<String>,
state: Arc<AtomicU8>,
started: Arc<AtomicBool>,
endpoint_client: Arc<Mutex<Option<moq_native_ietf::quic::Client>>>,
outbound_tracks: Arc<Mutex<Draft14TrackWriters>>,
tasks: Arc<Mutex<Vec<tokio::task::JoinHandle<()>>>>,
message_handler: Arc<RwLock<Option<NativeMoQMessageHandler>>>,
terminal_handler: Arc<RwLock<Option<NativeMoQTerminalHandler>>>,
pending_messages: Arc<Mutex<Vec<Bytes>>>,
}
impl NativeMoQSession {
pub async fn new(
local_node_id: &str,
remote_node_id: &str,
config: &crate::client::MoQConfig,
) -> Result<Self> {
let relay_url = config.relay_url.trim();
ensure!(!relay_url.is_empty(), "moq relay_url is required");
ensure!(
relay_url.starts_with("https://"),
"moq relay_url must use https/WebTransport"
);
ensure!(
!relay_url.contains('#'),
"moq relay_url must not contain a fragment"
);
ensure!(
!relay_url_has_jwt_query(relay_url),
"moq relay_url must not contain a jwt query parameter; use access_token"
);
let access_token = config
.access_token
.as_deref()
.map(str::trim)
.map(str::to_string);
if let Some(token) = access_token.as_deref() {
ensure!(!token.is_empty(), "moq access_token must not be empty");
ensure!(
token.len() <= MOQ_MAX_ACCESS_TOKEN_BYTES,
"moq access_token exceeds the maximum supported size"
);
ensure!(
token
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')),
"moq access_token must be a URL-safe JWT"
);
}
Ok(Self {
local_node_id: local_node_id.to_string(),
remote_node_id: remote_node_id.to_string(),
relay_url: relay_url.to_string(),
access_token,
state: Arc::new(AtomicU8::new(0)),
started: Arc::new(AtomicBool::new(false)),
endpoint_client: Arc::new(Mutex::new(None)),
outbound_tracks: Arc::new(Mutex::new(HashMap::new())),
tasks: Arc::new(Mutex::new(Vec::new())),
message_handler: Arc::new(RwLock::new(None)),
terminal_handler: Arc::new(RwLock::new(None)),
pending_messages: Arc::new(Mutex::new(Vec::new())),
})
}
pub fn state(&self) -> NativeMoQState {
decode_state(self.state.load(Ordering::SeqCst))
}
#[cfg(test)]
pub(crate) fn force_state_for_test(&self, state: NativeMoQState) {
self.state.store(encode_state(state), Ordering::SeqCst);
}
pub async fn start(&self) -> Result<()> {
if self.started.swap(true, Ordering::SeqCst) {
return Ok(());
}
self.state.store(1, Ordering::SeqCst);
if self.relay_url.contains(".invalid") {
self.state.store(3, Ordering::SeqCst);
bail!(
"moq relay is unreachable: {}",
relay_diagnostic_label(&self.relay_url)
);
}
let setup = tokio::time::timeout(MOQ_SETUP_TIMEOUT, self.connect_draft14()).await;
let (client, session, publisher, subscriber) = match setup {
Ok(Ok(connected)) => connected,
Ok(Err(error)) => {
self.state.store(3, Ordering::SeqCst);
return Err(error);
}
Err(_) => {
self.state.store(3, Ordering::SeqCst);
bail!(
"MoQ Draft 14 setup timed out after {} ms",
MOQ_SETUP_TIMEOUT.as_millis()
);
}
};
*self.endpoint_client.lock().await = Some(client);
self.state.store(2, Ordering::SeqCst);
self.spawn_session_driver(session).await;
self.spawn_publisher(publisher).await;
self.spawn_subscription(subscriber.clone(), self.direct_inbound_track_name())
.await;
self.spawn_subscription(subscriber, MOQ_BROADCAST_TRACK.to_string())
.await;
Ok(())
}
pub async fn send(&self, data: &[u8]) -> Result<()> {
ensure!(
self.state() == NativeMoQState::Connected,
"native moq session is not started"
);
ensure!(
data.len() <= MOQ_MAX_OBJECT_PAYLOAD_BYTES,
"native moq object payload too large: {}",
data.len()
);
let track_name = self.direct_outbound_track_name();
let mut tracks = self.outbound_tracks.lock().await;
let Some(track) = tracks.get_mut(&track_name) else {
bail!("native moq track is not subscribed yet; caller must fall back to iroh");
};
let mut subgroup = track
.append(0)
.context("failed to create MoQ Draft 14 subgroup")?;
subgroup
.write(Bytes::copy_from_slice(data))
.context("failed to write MoQ Draft 14 object")?;
Ok(())
}
pub async fn is_peer_data_subscribed(&self) -> bool {
self.outbound_tracks
.lock()
.await
.contains_key(&self.direct_outbound_track_name())
}
pub async fn wait_for_peer_data_subscription(&self, timeout: Duration) -> bool {
tokio::time::timeout(timeout, async {
loop {
if self.is_peer_data_subscribed().await {
return true;
}
if self.state() != NativeMoQState::Connected {
return false;
}
tokio::time::sleep(Duration::from_millis(25)).await;
}
})
.await
.unwrap_or(false)
}
pub async fn close_gracefully(&self) {
self.state.store(4, Ordering::SeqCst);
self.outbound_tracks.lock().await.clear();
let tasks = self.tasks.lock().await.drain(..).collect::<Vec<_>>();
for task in &tasks {
task.abort();
}
for task in tasks {
let _ = task.await;
}
self.endpoint_client.lock().await.take();
}
pub fn close(&self) {
self.state.store(4, Ordering::SeqCst);
let session = self.clone();
tokio::spawn(async move {
session.close_gracefully().await;
});
}
pub async fn set_message_handler(&self, handler: NativeMoQMessageHandler) {
*self.message_handler.write().await = Some(handler.clone());
let pending = {
let mut guard = self.pending_messages.lock().await;
std::mem::take(&mut *guard)
};
for message in pending {
handler(message).await;
}
}
pub async fn set_terminal_handler(&self, handler: NativeMoQTerminalHandler) {
*self.terminal_handler.write().await = Some(handler);
}
#[cfg(feature = "test-harness")]
pub async fn force_failure_for_test(&self) -> bool {
if !mark_failed_unless_closed(&self.state) {
return false;
}
self.outbound_tracks.lock().await.clear();
let tasks = self.tasks.lock().await.drain(..).collect::<Vec<_>>();
for task in &tasks {
task.abort();
}
for task in tasks {
let _ = task.await;
}
self.endpoint_client.lock().await.take();
self.notify_terminal_failure().await;
true
}
async fn connect_draft14(
&self,
) -> Result<(
moq_native_ietf::quic::Client,
moq_transport::session::Session,
moq_transport::session::Publisher,
moq_transport::session::Subscriber,
)> {
let connection_url = relay_connection_url(&self.relay_url, self.access_token.as_deref());
let url = connection_url
.parse::<url::Url>()
.context("invalid MoQ relay URL")?;
let tls = moq_native_ietf::tls::Args {
disable_verify: relay_url_uses_loopback_host(&self.relay_url),
..Default::default()
}
.load()
.context("failed to configure MoQ TLS")?;
let quic = moq_native_ietf::quic::Args {
tls: moq_native_ietf::tls::Args {
disable_verify: relay_url_uses_loopback_host(&self.relay_url),
..Default::default()
},
..Default::default()
}
.load()
.or_else(|_| {
moq_native_ietf::quic::Config::new(
"0.0.0.0:0".parse().expect("valid fallback bind address"),
None,
tls,
)
})
.context("failed to configure MoQ QUIC")?;
let endpoint = moq_native_ietf::quic::Endpoint::new(quic)
.context("failed to create MoQ Draft 14 endpoint")?;
let client = endpoint.client;
let relay_label = relay_diagnostic_label(&self.relay_url);
let (webtransport, _, transport) = client.connect(&url, None).await.map_err(|error| {
anyhow!(
"failed to connect MoQ Draft 14 relay {relay_label}: {}",
redact_access_token(&error.to_string(), self.access_token.as_deref())
)
})?;
let (session, publisher, subscriber) =
moq_transport::session::Session::connect(webtransport, None, transport)
.await
.context("MoQ Draft 14 SETUP exchange failed")?;
Ok((client, session, publisher, subscriber))
}
async fn spawn_session_driver(&self, session: moq_transport::session::Session) {
let state = self.state.clone();
let terminal_handler = self.terminal_handler.clone();
self.push_task(tokio::spawn(async move {
if let Err(error) = session.run().await {
eprintln!("[NativeMoQ] Draft 14 session failed: {error:#}");
}
if mark_failed_unless_closed(&state) {
notify_terminal_failure(&terminal_handler).await;
}
}))
.await;
}
async fn spawn_publisher(&self, mut publisher: moq_transport::session::Publisher) {
let namespace = TrackNamespace::from_utf8_path(&self.local_node_id);
let (_tracks_writer, mut requests, tracks_reader) = serve::Tracks::new(namespace).produce();
let allowed_direct = self.direct_outbound_track_name();
let outbound_tracks = self.outbound_tracks.clone();
let state = self.state.clone();
let terminal_handler = self.terminal_handler.clone();
self.push_task(tokio::spawn(async move {
let announce = publisher.announce(tracks_reader);
tokio::pin!(announce);
loop {
tokio::select! {
result = &mut announce => {
match result {
Ok(()) => {}
Err(error) => eprintln!("[NativeMoQ] Draft 14 announce failed: {error:#}"),
}
if mark_failed_unless_closed(&state) {
notify_terminal_failure(&terminal_handler).await;
}
return;
}
requested = requests.next() => {
let Some(track) = requested else {
if mark_failed_unless_closed(&state) {
notify_terminal_failure(&terminal_handler).await;
}
return;
};
let name = track.name.clone();
if name != allowed_direct && name != MOQ_BROADCAST_TRACK {
let _ = track.close(serve::ServeError::NotFound);
continue;
}
match track.subgroups() {
Ok(writer) => {
outbound_tracks.lock().await.insert(name, writer);
}
Err(error) => {
eprintln!("[NativeMoQ] failed to serve Draft 14 track: {error}");
}
}
}
}
}
}))
.await;
}
async fn spawn_subscription(
&self,
subscriber: moq_transport::session::Subscriber,
track_name: String,
) {
let namespace = TrackNamespace::from_utf8_path(&self.remote_node_id);
let state = self.state.clone();
let message_handler = self.message_handler.clone();
let pending_messages = self.pending_messages.clone();
self.push_task(tokio::spawn(async move {
while decode_state(state.load(Ordering::SeqCst)) == NativeMoQState::Connected {
let (track_writer, track_reader) =
serve::Track::new(namespace.clone(), track_name.clone()).produce();
let mut subscriber = subscriber.clone();
match subscriber.subscribe_open(track_writer).await {
Ok(subscription) => {
let result = receive_track(
track_reader,
message_handler.clone(),
pending_messages.clone(),
)
.await;
drop(subscription);
if let Err(error) = result {
eprintln!(
"[NativeMoQ] Draft 14 subscription {track_name} ended: {error:#}"
);
}
}
Err(error) => {
eprintln!(
"[NativeMoQ] Draft 14 subscription {track_name} not ready: {error}"
);
}
}
if decode_state(state.load(Ordering::SeqCst)) != NativeMoQState::Connected {
return;
}
tokio::time::sleep(MOQ_SUBSCRIBE_RETRY_DELAY).await;
}
}))
.await;
}
async fn push_task(&self, task: tokio::task::JoinHandle<()>) {
self.tasks.lock().await.push(task);
}
async fn notify_terminal_failure(&self) {
notify_terminal_failure(&self.terminal_handler).await;
}
fn direct_outbound_track_name(&self) -> String {
format!("data / {} ", self.remote_node_id)
}
fn direct_inbound_track_name(&self) -> String {
format!("data / {} ", self.local_node_id)
}
}
async fn receive_track(
track: serve::TrackReader,
message_handler: Arc<RwLock<Option<NativeMoQMessageHandler>>>,
pending_messages: Arc<Mutex<Vec<Bytes>>>,
) -> Result<()> {
let TrackReaderMode::Subgroups(mut groups) = track.mode().await? else {
bail!("MoQ Draft 14 peer selected an unsupported track mode");
};
while let Some(mut group) = groups.next().await? {
while let Some(payload) = group.read_next().await? {
ensure!(
payload.len() <= MOQ_MAX_OBJECT_PAYLOAD_BYTES,
"native moq object payload too large: {}",
payload.len()
);
dispatch_or_buffer_message(&message_handler, &pending_messages, payload).await;
}
}
Ok(())
}
async fn dispatch_or_buffer_message(
message_handler: &Arc<RwLock<Option<NativeMoQMessageHandler>>>,
pending_messages: &Arc<Mutex<Vec<Bytes>>>,
payload: Bytes,
) {
let handler = message_handler.read().await.clone();
if let Some(handler) = handler {
handler(payload).await;
} else {
pending_messages.lock().await.push(payload);
}
}
fn mark_failed_unless_closed(state: &AtomicU8) -> bool {
loop {
let current = state.load(Ordering::SeqCst);
match decode_state(current) {
NativeMoQState::Closed | NativeMoQState::Failed => return false,
NativeMoQState::Idle | NativeMoQState::Connecting | NativeMoQState::Connected => {
if state
.compare_exchange(current, 3, Ordering::SeqCst, Ordering::SeqCst)
.is_ok()
{
return true;
}
}
}
}
}
async fn notify_terminal_failure(terminal_handler: &Arc<RwLock<Option<NativeMoQTerminalHandler>>>) {
if let Some(handler) = terminal_handler.read().await.clone() {
handler().await;
}
}
#[cfg(test)]
fn encode_state(state: NativeMoQState) -> u8 {
match state {
NativeMoQState::Idle => 0,
NativeMoQState::Connecting => 1,
NativeMoQState::Connected => 2,
NativeMoQState::Failed => 3,
NativeMoQState::Closed => 4,
}
}
fn decode_state(state: u8) -> NativeMoQState {
match state {
1 => NativeMoQState::Connecting,
2 => NativeMoQState::Connected,
3 => NativeMoQState::Failed,
4 => NativeMoQState::Closed,
_ => NativeMoQState::Idle,
}
}
fn relay_url_uses_loopback_host(relay_url: &str) -> bool {
relay_url.starts_with("https://localhost")
|| relay_url.starts_with("https://127.")
|| relay_url.starts_with("https://[::1]")
}
fn relay_url_has_jwt_query(relay_url: &str) -> bool {
relay_url
.split_once('?')
.map(|(_, query)| {
query
.split('&')
.map(|pair| pair.split_once('=').map_or(pair, |(key, _)| key))
.any(|key| key.eq_ignore_ascii_case("jwt"))
})
.unwrap_or(false)
}
fn relay_connection_url(relay_url: &str, access_token: Option<&str>) -> String {
let Some(token) = access_token else {
return relay_url.to_string();
};
let separator = if relay_url.contains('?') { '&' } else { '?' };
format!("{relay_url}{separator}jwt={token}")
}
fn relay_diagnostic_label(relay_url: &str) -> &str {
relay_url
.split_once('?')
.map_or(relay_url, |(base, _)| base)
}
fn redact_access_token(value: &str, access_token: Option<&str>) -> String {
let Some(token) = access_token else {
return value.to_string();
};
value.replace(token, "[redacted]")
}
#[cfg(test)]
mod tests {
use super::*;
use crate::client::MoQConfig;
fn make_config(relay_url: &str) -> MoQConfig {
MoQConfig {
relay_url: relay_url.to_string(),
access_token: None,
}
}
#[tokio::test]
async fn invalid_relay_url_reaches_failed() {
let session =
NativeMoQSession::new("local", "remote", &make_config("https://relay.invalid/moq"))
.await
.unwrap();
assert!(session.start().await.is_err());
assert_eq!(session.state(), NativeMoQState::Failed);
}
#[tokio::test]
async fn send_requires_connected_and_subscribed_route() {
let session = NativeMoQSession::new(
"local",
"remote",
&make_config("https://relay.example.com/moq"),
)
.await
.unwrap();
let error = session.send(b"hello").await.unwrap_err();
assert!(error.to_string().contains("not started"));
session.force_state_for_test(NativeMoQState::Connected);
let error = session.send(b"hello").await.unwrap_err();
assert!(error.to_string().contains("not subscribed"));
}
#[tokio::test]
async fn close_reaches_closed_state() {
let session = NativeMoQSession::new(
"local",
"remote",
&make_config("https://relay.example.com/moq"),
)
.await
.unwrap();
session.close_gracefully().await;
assert_eq!(session.state(), NativeMoQState::Closed);
}
#[cfg(feature = "test-harness")]
#[tokio::test]
async fn draft14_local_relay_carries_bilateral_native_payloads() {
let probe = std::net::UdpSocket::bind("127.0.0.1:0").unwrap();
let port = probe.local_addr().unwrap().port();
drop(probe);
let relay = tokio::spawn(async move {
let _ = test_relay::run(port).await;
});
tokio::time::sleep(Duration::from_millis(100)).await;
let config = make_config(&format!("https://localhost:{port}/moq"));
let left = NativeMoQSession::new("left", "right", &config)
.await
.unwrap();
let right = NativeMoQSession::new("right", "left", &config)
.await
.unwrap();
left.start().await.unwrap();
right.start().await.unwrap();
assert!(
left.wait_for_peer_data_subscription(Duration::from_secs(5))
.await
);
assert!(
right
.wait_for_peer_data_subscription(Duration::from_secs(5))
.await
);
let (left_tx, left_rx) = tokio::sync::oneshot::channel::<Bytes>();
let left_tx = Arc::new(Mutex::new(Some(left_tx)));
left.set_message_handler(Arc::new(move |payload| {
let left_tx = left_tx.clone();
Box::pin(async move {
if let Some(tx) = left_tx.lock().await.take() {
let _ = tx.send(payload);
}
})
}))
.await;
let (right_tx, right_rx) = tokio::sync::oneshot::channel::<Bytes>();
let right_tx = Arc::new(Mutex::new(Some(right_tx)));
right
.set_message_handler(Arc::new(move |payload| {
let right_tx = right_tx.clone();
Box::pin(async move {
if let Some(tx) = right_tx.lock().await.take() {
let _ = tx.send(payload);
}
})
}))
.await;
left.send(b"left-to-right").await.unwrap();
right.send(b"right-to-left").await.unwrap();
assert_eq!(
tokio::time::timeout(Duration::from_secs(3), right_rx)
.await
.unwrap()
.unwrap(),
Bytes::from_static(b"left-to-right")
);
assert_eq!(
tokio::time::timeout(Duration::from_secs(3), left_rx)
.await
.unwrap()
.unwrap(),
Bytes::from_static(b"right-to-left")
);
left.close_gracefully().await;
right.close_gracefully().await;
relay.abort();
}
#[cfg(feature = "test-harness")]
#[tokio::test]
async fn draft14_replacement_releases_the_previous_node_namespace() {
let probe = std::net::UdpSocket::bind("127.0.0.1:0").unwrap();
let port = probe.local_addr().unwrap().port();
drop(probe);
let relay = tokio::spawn(async move {
let _ = test_relay::run(port).await;
});
tokio::time::sleep(Duration::from_millis(100)).await;
let config = make_config(&format!("https://localhost:{port}/moq"));
let first_left = NativeMoQSession::new("left", "right", &config)
.await
.unwrap();
let right = NativeMoQSession::new("right", "left", &config)
.await
.unwrap();
first_left.start().await.unwrap();
right.start().await.unwrap();
assert!(
first_left
.wait_for_peer_data_subscription(Duration::from_secs(5))
.await
);
assert!(
right
.wait_for_peer_data_subscription(Duration::from_secs(5))
.await
);
first_left.close_gracefully().await;
let replacement_left = NativeMoQSession::new("left", "right", &config)
.await
.unwrap();
replacement_left.start().await.unwrap();
assert!(
replacement_left
.wait_for_peer_data_subscription(Duration::from_secs(5))
.await,
"replacement publisher never became data-ready"
);
assert!(
right
.wait_for_peer_data_subscription(Duration::from_secs(5))
.await,
"existing peer did not resubscribe to the replacement publisher"
);
let (received_tx, received_rx) = tokio::sync::oneshot::channel::<Bytes>();
let received_tx = Arc::new(Mutex::new(Some(received_tx)));
right
.set_message_handler(Arc::new(move |payload| {
let received_tx = received_tx.clone();
Box::pin(async move {
if let Some(tx) = received_tx.lock().await.take() {
let _ = tx.send(payload);
}
})
}))
.await;
replacement_left.send(b"replacement-payload").await.unwrap();
assert_eq!(
tokio::time::timeout(Duration::from_secs(3), received_rx)
.await
.unwrap()
.unwrap(),
Bytes::from_static(b"replacement-payload")
);
replacement_left.close_gracefully().await;
right.close_gracefully().await;
relay.abort();
}
#[test]
fn access_token_is_connection_only_and_redacted() {
let token = "header.payload.signature";
let relay_url = "https://relay.example.com/moq?existing=value";
assert_eq!(
relay_connection_url(relay_url, Some(token)),
format!("{relay_url}&jwt={token}")
);
assert_eq!(
relay_diagnostic_label(relay_url),
"https://relay.example.com/moq"
);
assert_eq!(
redact_access_token(
&format!("failed to connect https://relay.example.com/moq?jwt={token}"),
Some(token),
),
"failed to connect https://relay.example.com/moq?jwt=[redacted]"
);
}
#[tokio::test]
async fn inline_jwt_and_non_url_safe_tokens_are_rejected() {
let inline = make_config("https://relay.example.com/moq?jwt=secret");
let error = NativeMoQSession::new("local", "remote", &inline)
.await
.err()
.expect("inline JWT must be rejected");
assert!(error.to_string().contains("use access_token"));
let invalid = MoQConfig {
relay_url: "https://relay.example.com/moq".to_string(),
access_token: Some("not a jwt&leak".to_string()),
};
let error = NativeMoQSession::new("local", "remote", &invalid)
.await
.err()
.expect("unsafe token must be rejected");
assert!(error.to_string().contains("URL-safe JWT"));
}
}