use crate::{
coding::{KeyValuePairs, Location, ReasonPhrase, TrackName, TrackNamespace},
data,
message::{self, RequestErrorCode},
serve::{self, ServeError, TrackReader, TrackWriterMode},
watch::State,
};
use super::Subscriber;
pub(crate) struct PublishReceivedState {
done: bool,
closed: Result<(), ServeError>,
}
impl Default for PublishReceivedState {
fn default() -> Self {
Self {
done: false,
closed: Ok(()),
}
}
}
pub struct PublishReceived {
session: Subscriber,
state: State<PublishReceivedState>,
reader: Option<TrackReader>,
request_id: u64,
track_alias: u64,
namespace: TrackNamespace,
name: TrackName,
initial_forward: bool,
largest_location: Option<Location>,
ok: bool,
error: Option<ServeError>,
}
impl PublishReceived {
#[allow(clippy::too_many_arguments)]
pub(super) fn new(
session: Subscriber,
request_id: u64,
track_alias: u64,
namespace: TrackNamespace,
name: TrackName,
initial_forward: bool,
largest_location: Option<Location>,
reader: TrackReader,
state: State<PublishReceivedState>,
) -> Self {
Self {
session,
state,
reader: Some(reader),
request_id,
track_alias,
namespace,
name,
initial_forward,
largest_location,
ok: false,
error: None,
}
}
pub fn take_reader(&mut self) -> Result<TrackReader, ServeError> {
self.reader.take().ok_or(ServeError::Done)
}
pub fn accept(&mut self, forward: bool) -> Result<(), ServeError> {
if self.ok {
return Err(ServeError::Duplicate);
}
let mut params = KeyValuePairs::default();
params.set_forward(forward);
self.session.send_message(message::PublishOk {
id: self.request_id,
params,
});
self.ok = true;
Ok(())
}
pub fn ok(&mut self, forward: bool) -> Result<TrackReader, ServeError> {
let reader = self.take_reader()?;
self.accept(forward)?;
Ok(reader)
}
pub fn close(mut self, err: ServeError) {
self.error = Some(err);
}
pub async fn closed(&self) -> Result<(), ServeError> {
loop {
{
let state = self.state.lock();
match state.closed.clone() {
Ok(()) => {}
Err(ServeError::Done) => return Ok(()),
Err(err) => return Err(err),
}
match state.modified() {
Some(notify) => notify,
None => return Ok(()),
}
}
.await;
}
}
pub fn namespace(&self) -> &TrackNamespace {
&self.namespace
}
pub fn name(&self) -> &TrackName {
&self.name
}
pub fn track_alias(&self) -> u64 {
self.track_alias
}
pub fn initial_forward(&self) -> bool {
self.initial_forward
}
pub fn largest_location(&self) -> Option<Location> {
self.largest_location
}
}
impl Drop for PublishReceived {
fn drop(&mut self) {
if self.ok {
return;
}
let err = self.error.clone().unwrap_or(ServeError::Cancel);
let error_code = match &err {
ServeError::Cancel | ServeError::Done => RequestErrorCode::Uninterested as u64,
ServeError::Duplicate => RequestErrorCode::DuplicateSubscription as u64,
ServeError::NotFound | ServeError::NotFoundWithId(_, _) => {
RequestErrorCode::DoesNotExist as u64
}
ServeError::NotImplemented(_) | ServeError::NotImplementedWithId(_, _) => {
RequestErrorCode::NotSupported as u64
}
ServeError::Internal(_) | ServeError::InternalWithId(_, _) => {
RequestErrorCode::InternalError as u64
}
ServeError::Closed(code) => *code,
_ => RequestErrorCode::InternalError as u64,
};
self.session.send_request_error(
"publish",
message::RequestError {
id: self.request_id,
error_code,
retry_interval: 0,
reason: ReasonPhrase("uninterested".to_string()),
},
);
self.session.remove_publish_received(self.request_id);
}
}
pub(crate) struct PublishReceivedRecv {
state: State<PublishReceivedState>,
writer: Option<TrackWriterMode>,
}
impl PublishReceivedRecv {
#[allow(clippy::too_many_arguments)]
pub(super) fn produce(
session: Subscriber,
request_id: u64,
track_alias: u64,
namespace: TrackNamespace,
name: TrackName,
initial_forward: bool,
largest_location: Option<Location>,
writer: serve::TrackWriter,
reader: TrackReader,
) -> (PublishReceived, PublishReceivedRecv) {
let (app_state, transport_state) = State::<PublishReceivedState>::default().split();
let app = PublishReceived::new(
session,
request_id,
track_alias,
namespace,
name,
initial_forward,
largest_location,
reader,
app_state,
);
let recv = Self {
state: transport_state,
writer: Some(writer.into()),
};
(app, recv)
}
pub fn subgroup(
&mut self,
header: data::SubgroupHeader,
) -> Result<serve::SubgroupWriter, ServeError> {
let writer = self.writer.take().ok_or(ServeError::Done)?;
let mut subgroups = match writer {
TrackWriterMode::Track(track) => track.subgroups()?,
TrackWriterMode::Subgroups(subgroups) => subgroups,
_ => return Err(ServeError::Mode),
};
let subgroup_writer = subgroups.create(serve::Subgroup {
group_id: header.group_id,
subgroup_id: header.subgroup_id.unwrap_or(0),
priority: header.publisher_priority,
})?;
self.writer = Some(subgroups.into());
Ok(subgroup_writer)
}
pub fn datagram(&mut self, datagram: data::Datagram) -> Result<(), ServeError> {
let writer = self.writer.take().ok_or(ServeError::Done)?;
match writer {
TrackWriterMode::Track(track) => {
let mut datagrams = track.datagrams()?;
datagrams.write(serve::Datagram {
group_id: datagram.group_id,
object_id: datagram.object_id.unwrap_or(0),
priority: datagram.publisher_priority,
payload: datagram.payload.unwrap_or_default(),
extension_headers: datagram.extension_headers.unwrap_or_default(),
})?;
self.writer = Some(TrackWriterMode::Datagrams(datagrams));
Ok(())
}
TrackWriterMode::Datagrams(mut datagrams) => {
datagrams.write(serve::Datagram {
group_id: datagram.group_id,
object_id: datagram.object_id.unwrap_or(0),
priority: datagram.publisher_priority,
payload: datagram.payload.unwrap_or_default(),
extension_headers: datagram.extension_headers.unwrap_or_default(),
})?;
self.writer = Some(TrackWriterMode::Datagrams(datagrams));
Ok(())
}
other => {
self.writer = Some(other);
Err(ServeError::Mode)
}
}
}
pub fn recv_done(&mut self, status_code: u64) {
if let Some(mut state) = self.state.lock_mut() {
state.done = true;
state.closed = if status_code == message::PublishDoneCode::TrackEnded as u64 {
Err(ServeError::Done)
} else {
Err(ServeError::Closed(status_code))
};
}
self.writer = None;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
coding::TrackNamespace,
serve::Track,
session::{Queue, RequestId, SessionId},
};
fn make_pair(
request_id: u64,
) -> (
PublishReceived,
PublishReceivedRecv,
crate::session::Subscriber,
) {
let rid = RequestId::new(0, 100, 100, 0);
let subscriber = crate::session::Subscriber::new(
Queue::default(),
Queue::default(),
None,
rid,
crate::session::PendingRequests::default(),
SessionId::generate(),
);
let (writer, reader) =
Track::new(TrackNamespace::from_utf8_path("test"), "0.mp4").produce();
let (pr, recv) = PublishReceivedRecv::produce(
subscriber.clone(),
request_id,
42,
TrackNamespace::from_utf8_path("test"),
"0.mp4".into(),
true,
None,
writer,
reader,
);
(pr, recv, subscriber)
}
#[test]
fn take_reader_returns_once() {
let (mut pr, _recv, _sub) = make_pair(0);
assert!(pr.take_reader().is_ok());
assert!(
pr.take_reader().is_err(),
"reader must only be given out once"
);
}
#[test]
fn recv_done_closes_writer_and_sets_state() {
let (_pr, mut recv, _sub) = make_pair(1);
assert!(recv.writer.is_some());
recv.recv_done(message::PublishDoneCode::TrackEnded as u64);
assert!(
recv.writer.is_none(),
"writer must be dropped after PUBLISH_DONE"
);
}
#[test]
fn recv_done_non_track_ended_stores_code() {
let (_pr, mut recv, _sub) = make_pair(2);
recv.recv_done(message::PublishDoneCode::Expired as u64);
assert!(recv.writer.is_none());
}
#[tokio::test]
async fn closed_returns_ok_for_track_ended() {
let (pr, mut recv, _sub) = make_pair(3);
recv.recv_done(message::PublishDoneCode::TrackEnded as u64);
assert_eq!(pr.closed().await, Ok(()));
}
#[tokio::test]
async fn closed_returns_error_for_non_track_ended() {
let (pr, mut recv, _sub) = make_pair(4);
recv.recv_done(message::PublishDoneCode::Expired as u64);
assert!(matches!(
pr.closed().await,
Err(ServeError::Closed(code)) if code == message::PublishDoneCode::Expired as u64
));
}
#[test]
fn publish_received_drop_without_ok_does_not_panic() {
let (pr, _recv, _sub) = make_pair(5);
drop(pr);
}
}