use std::marker::PhantomData;
use std::ops::{Deref, DerefMut};
use std::sync::{Arc, Mutex, MutexGuard};
use serde::Serialize;
use serde::de::DeserializeOwned;
use super::{Encoded, Encoder};
use crate::{Error, Result};
pub use super::Config;
fn take<T>(inner: &Mutex<Inner<T>>) -> MutexGuard<'_, Inner<T>> {
inner.lock().unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub struct Producer<T> {
inner: Arc<Mutex<Inner<T>>>,
_marker: PhantomData<fn(T)>,
}
impl<T> Clone for Producer<T> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
_marker: PhantomData,
}
}
}
impl<T> Producer<T> {
pub fn consume(&self) -> moq_net::track::Subscriber {
take(&self.inner).track.inner.subscribe(None)
}
pub fn is_used(&self) -> bool {
take(&self.inner).track.inner.is_used()
}
pub fn demand(&self) -> moq_net::track::Demand {
take(&self.inner).track.inner.demand()
}
}
impl<T: Serialize> Producer<T> {
pub fn new(track: moq_net::track::Producer, config: Config) -> Self {
Self {
inner: Arc::new(Mutex::new(Inner {
track: Track {
inner: track,
group: None,
deltas: config.delta_ratio != 0,
},
encoder: Encoder::new(config),
aborted: None,
})),
_marker: PhantomData,
}
}
pub fn update(&mut self, value: &T) -> Result<()> {
take(&self.inner).update(value)
}
pub fn mutate(&mut self, f: impl FnOnce(&mut T)) -> Result<()>
where
T: Default + DeserializeOwned,
{
let mut guard = self.modify()?;
f(&mut guard);
guard.commit()
}
pub fn modify(&mut self) -> Result<Guard<'_, T>>
where
T: Default + DeserializeOwned,
{
let inner = take(&self.inner);
inner.open()?;
let value = match inner.encoder.value() {
Some(last) => serde_json::from_value(last.clone())?,
None => T::default(),
};
Ok(Guard {
inner,
value,
dirty: false,
})
}
pub fn cut(&mut self) -> Result<()> {
take(&self.inner).cut()
}
pub fn finish(&mut self) -> Result<()> {
take(&self.inner).finish()
}
pub fn abort(self, err: moq_net::Error) -> Result<()> {
let mut inner = take(&self.inner);
inner.abort(err.into());
Ok(())
}
}
pub struct Guard<'a, T: Serialize> {
inner: MutexGuard<'a, Inner<T>>,
value: T,
dirty: bool,
}
impl<T: Serialize> Guard<'_, T> {
pub fn commit(mut self) -> Result<()> {
self.publish()
}
fn publish(&mut self) -> Result<()> {
if !self.dirty {
return Ok(());
}
self.dirty = false;
self.inner.update(&self.value)
}
}
impl<T: Serialize> Deref for Guard<'_, T> {
type Target = T;
fn deref(&self) -> &T {
&self.value
}
}
impl<T: Serialize> DerefMut for Guard<'_, T> {
fn deref_mut(&mut self) -> &mut T {
self.dirty = true;
&mut self.value
}
}
impl<T: Serialize> Drop for Guard<'_, T> {
fn drop(&mut self) {
if std::thread::panicking() {
return;
}
if let Err(err) = self.publish() {
tracing::error!(%err, "failed to publish JSON value on guard drop, aborting the track");
self.inner.abort(err);
}
}
}
struct Inner<T> {
track: Track,
encoder: Encoder<T>,
aborted: Option<Error>,
}
impl<T> Inner<T> {
fn open(&self) -> Result<()> {
if let Some(err) = &self.aborted {
return Err(err.clone());
}
let track = &self.track.inner;
let next = track.latest().map_or(0, |latest| latest.saturating_add(1));
if track.is_closed() || track.final_sequence().is_some_and(|fin| next >= fin) {
return Err(moq_net::Error::Closed.into());
}
Ok(())
}
fn abort(&mut self, err: Error) {
let reason = match &err {
Error::Net(err) => err.clone(),
_ => moq_net::StreamError::Internal.into(),
};
self.encoder.reset();
self.track.group = None;
let _ = self.track.inner.clone().abort(reason);
self.aborted = Some(err);
}
fn cut(&mut self) -> Result<()> {
if self.track.group.is_none() {
return Ok(());
}
self.encoder.reset();
self.track.cut()
}
}
impl<T: Serialize> Inner<T> {
fn update(&mut self, value: &T) -> Result<()> {
let Inner { track, encoder, .. } = self;
let Some(frame) = encoder.update(value)? else {
return Ok(());
};
track.write(&frame)?;
frame.commit();
Ok(())
}
fn finish(&mut self) -> Result<()> {
self.encoder.reset();
self.track.finish()
}
}
struct Track {
inner: moq_net::track::Producer,
group: Option<moq_net::group::Producer>,
deltas: bool,
}
impl Track {
fn cut(&mut self) -> Result<()> {
if let Some(group) = self.group.take() {
group.finish()?;
}
Ok(())
}
fn write(&mut self, encoded: &Encoded) -> Result<()> {
if encoded.payload.len() as u64 > moq_net::group::MAX_CACHE_BYTES {
return Err(moq_net::Error::FrameTooLarge.into());
}
match encoded.keyframe {
true => self.write_snapshot(encoded.payload.clone()),
false => self.write_delta(encoded.payload.clone()),
}
}
fn write_snapshot(&mut self, payload: bytes::Bytes) -> Result<()> {
if let Some(group) = self.group.take() {
group.finish()?;
}
let mut group = self.inner.append_group()?;
if let Err(err) = group.write_frame(moq_net::Timestamp::now(), payload) {
let _ = group.finish();
return Err(err.into());
}
match self.deltas {
true => self.group = Some(group),
false => group.finish()?,
}
Ok(())
}
fn write_delta(&mut self, payload: bytes::Bytes) -> Result<()> {
self.group
.as_mut()
.expect("the encoder only emits a delta after a snapshot opened a group")
.write_frame(moq_net::Timestamp::now(), payload)?;
Ok(())
}
fn finish(&mut self) -> Result<()> {
if let Some(group) = self.group.take() {
group.finish()?;
}
self.inner.finish()?;
Ok(())
}
}