moqtail 0.14.0

Draft 16-compliant Media-over-QUIC (MoQ) protocol library.
Documentation
// Copyright 2025 The MOQtail Authors
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
//     http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.

use crate::model::common::varint::BufVarIntExt;
use crate::model::error::ParseError;
use bytes::{Buf, Bytes};

use super::{
  client_setup::ClientSetup, constant::ControlMessageType, fetch::Fetch, fetch_cancel::FetchCancel,
  fetch_ok::FetchOk, goaway::GoAway, max_request_id::MaxRequestId, namespace::Namespace,
  namespace_done::NamespaceDone, publish::Publish, publish_done::PublishDone,
  publish_namespace::PublishNamespace, publish_namespace_cancel::PublishNamespaceCancel,
  publish_namespace_done::PublishNamespaceDone, publish_ok::PublishOk, request_error::RequestError,
  request_ok::RequestOk, request_update::RequestUpdate, requests_blocked::RequestsBlocked,
  server_setup::ServerSetup, subscribe::Subscribe, subscribe_namespace::SubscribeNamespace,
  subscribe_ok::SubscribeOk, switch::Switch, track_status::TrackStatus, unsubscribe::Unsubscribe,
  unsubscribe_namespace::UnsubscribeNamespace,
};

#[derive(Debug, Clone, PartialEq)]
pub enum ControlMessage {
  Namespace(Box<Namespace>),
  NamespaceDone(Box<NamespaceDone>),
  PublishNamespace(Box<PublishNamespace>),
  PublishNamespaceCancel(Box<PublishNamespaceCancel>),
  RequestOk(Box<RequestOk>),
  Publish(Box<Publish>),
  PublishOk(Box<PublishOk>),
  PublishDone(Box<PublishDone>),
  ClientSetup(Box<ClientSetup>),
  Fetch(Box<Fetch>),
  FetchCancel(Box<FetchCancel>),
  FetchOk(Box<FetchOk>),
  Goaway(Box<GoAway>),
  MaxRequestId(Box<MaxRequestId>),
  ServerSetup(Box<ServerSetup>),
  Subscribe(Box<Subscribe>),
  SubscribeOk(Box<SubscribeOk>),
  RequestUpdate(Box<RequestUpdate>),
  RequestsBlocked(Box<RequestsBlocked>),
  TrackStatus(Box<TrackStatus>),
  PublishNamespaceDone(Box<PublishNamespaceDone>),
  Unsubscribe(Box<Unsubscribe>),
  SubscribeNamespace(Box<SubscribeNamespace>),
  RequestError(Box<RequestError>),
  UnsubscribeNamespace(Box<UnsubscribeNamespace>),
  Switch(Box<Switch>),
}

pub trait ControlMessageTrait: std::fmt::Debug {
  fn serialize(&self) -> Result<Bytes, ParseError>;
  fn parse_payload(payload: &mut Bytes) -> Result<Box<Self>, ParseError>
  where
    Self: Sized;
  fn get_type(&self) -> ControlMessageType;
}

impl ControlMessage {
  pub fn deserialize(bytes: &mut Bytes) -> Result<Self, ParseError> {
    let message_type = bytes.get_vi()?;
    let msg_type = ControlMessageType::try_from(message_type)?;

    if bytes.remaining() < 2 {
      return Err(ParseError::NotEnoughBytes {
        context: "ControlMessage::deserialize(payload_length)",
        needed: 2,
        available: 0,
      });
    }
    let payload_length = bytes.get_u16() as usize;

    if bytes.remaining() < payload_length {
      return Err(ParseError::NotEnoughBytes {
        context: "ControlMessage::deserialize(payload_length)",
        needed: payload_length,
        available: bytes.remaining(),
      });
    }

    let mut payload = bytes.copy_to_bytes(payload_length);
    let message = match msg_type {
      ControlMessageType::Namespace => {
        Namespace::parse_payload(&mut payload).map(ControlMessage::Namespace)
      }
      ControlMessageType::NamespaceDone => {
        NamespaceDone::parse_payload(&mut payload).map(ControlMessage::NamespaceDone)
      }
      ControlMessageType::PublishNamespace => {
        PublishNamespace::parse_payload(&mut payload).map(ControlMessage::PublishNamespace)
      }
      ControlMessageType::PublishNamespaceCancel => {
        PublishNamespaceCancel::parse_payload(&mut payload)
          .map(ControlMessage::PublishNamespaceCancel)
      }
      ControlMessageType::PublishNamespaceDone => {
        PublishNamespaceDone::parse_payload(&mut payload).map(ControlMessage::PublishNamespaceDone)
      }
      ControlMessageType::RequestError => {
        RequestError::parse_payload(&mut payload).map(ControlMessage::RequestError)
      }
      ControlMessageType::RequestOk => {
        RequestOk::parse_payload(&mut payload).map(ControlMessage::RequestOk)
      }
      ControlMessageType::Publish => {
        Publish::parse_payload(&mut payload).map(ControlMessage::Publish)
      }
      ControlMessageType::PublishOk => {
        PublishOk::parse_payload(&mut payload).map(ControlMessage::PublishOk)
      }
      ControlMessageType::PublishDone => {
        PublishDone::parse_payload(&mut payload).map(ControlMessage::PublishDone)
      }
      ControlMessageType::ClientSetup => {
        ClientSetup::parse_payload(&mut payload).map(ControlMessage::ClientSetup)
      }
      ControlMessageType::Fetch => Fetch::parse_payload(&mut payload).map(ControlMessage::Fetch),
      ControlMessageType::FetchCancel => {
        FetchCancel::parse_payload(&mut payload).map(ControlMessage::FetchCancel)
      }
      ControlMessageType::FetchOk => {
        FetchOk::parse_payload(&mut payload).map(ControlMessage::FetchOk)
      }
      ControlMessageType::GoAway => GoAway::parse_payload(&mut payload).map(ControlMessage::Goaway),
      ControlMessageType::MaxRequestId => {
        MaxRequestId::parse_payload(&mut payload).map(ControlMessage::MaxRequestId)
      }
      ControlMessageType::ServerSetup => {
        ServerSetup::parse_payload(&mut payload).map(ControlMessage::ServerSetup)
      }
      ControlMessageType::Subscribe => {
        Subscribe::parse_payload(&mut payload).map(ControlMessage::Subscribe)
      }
      ControlMessageType::SubscribeOk => {
        SubscribeOk::parse_payload(&mut payload).map(ControlMessage::SubscribeOk)
      }
      ControlMessageType::RequestUpdate => {
        RequestUpdate::parse_payload(&mut payload).map(ControlMessage::RequestUpdate)
      }
      ControlMessageType::RequestsBlocked => {
        RequestsBlocked::parse_payload(&mut payload).map(ControlMessage::RequestsBlocked)
      }
      ControlMessageType::TrackStatus => {
        TrackStatus::parse_payload(&mut payload).map(ControlMessage::TrackStatus)
      }
      ControlMessageType::Unsubscribe => {
        Unsubscribe::parse_payload(&mut payload).map(ControlMessage::Unsubscribe)
      }
      ControlMessageType::SubscribeNamespace => {
        SubscribeNamespace::parse_payload(&mut payload).map(ControlMessage::SubscribeNamespace)
      }
      ControlMessageType::UnsubscribeNamespace => {
        UnsubscribeNamespace::parse_payload(&mut payload).map(ControlMessage::UnsubscribeNamespace)
      }
      ControlMessageType::Switch => Switch::parse_payload(&mut payload).map(ControlMessage::Switch),
    }
    .map_err(|err| ParseError::ProtocolViolation {
      context: "ControlMessage::deserialize(payload)",
      details: err.to_string(),
    })?;

    if payload.has_remaining() {
      return Err(ParseError::ProtocolViolation {
        context: "ControlMessage::deserialize(final_check)",
        details: format!(
          "Extra {} bytes remaining in payload after parsing",
          payload.remaining()
        ),
      });
    };
    Ok(message)
  }

  pub fn serialize(&self) -> Result<Bytes, ParseError> {
    match self {
      ControlMessage::Namespace(msg) => msg.serialize(),
      ControlMessage::NamespaceDone(msg) => msg.serialize(),
      ControlMessage::PublishNamespace(msg) => msg.serialize(),
      ControlMessage::PublishNamespaceCancel(msg) => msg.serialize(),
      ControlMessage::PublishNamespaceDone(msg) => msg.serialize(),
      ControlMessage::RequestError(msg) => msg.serialize(),
      ControlMessage::RequestOk(msg) => msg.serialize(),
      ControlMessage::Publish(msg) => msg.serialize(),
      ControlMessage::PublishOk(msg) => msg.serialize(),
      ControlMessage::PublishDone(msg) => msg.serialize(),
      ControlMessage::ClientSetup(msg) => msg.serialize(),
      ControlMessage::Fetch(msg) => msg.serialize(),
      ControlMessage::FetchCancel(msg) => msg.serialize(),
      ControlMessage::FetchOk(msg) => msg.serialize(),
      ControlMessage::Goaway(msg) => msg.serialize(),
      ControlMessage::MaxRequestId(msg) => msg.serialize(),
      ControlMessage::ServerSetup(msg) => msg.serialize(),
      ControlMessage::Subscribe(msg) => msg.serialize(),
      ControlMessage::SubscribeOk(msg) => msg.serialize(),
      ControlMessage::RequestUpdate(msg) => msg.serialize(),
      ControlMessage::RequestsBlocked(msg) => msg.serialize(),
      ControlMessage::TrackStatus(msg) => msg.serialize(),
      ControlMessage::Unsubscribe(msg) => msg.serialize(),
      ControlMessage::SubscribeNamespace(msg) => msg.serialize(),
      ControlMessage::UnsubscribeNamespace(msg) => msg.serialize(),
      ControlMessage::Switch(msg) => msg.serialize(),
    }
  }

  /// Returns the message type of the control message.
  pub fn get_type(&self) -> ControlMessageType {
    match self {
      ControlMessage::Namespace(_) => ControlMessageType::Namespace,
      ControlMessage::NamespaceDone(_) => ControlMessageType::NamespaceDone,
      ControlMessage::PublishNamespace(_) => ControlMessageType::PublishNamespace,
      ControlMessage::PublishNamespaceCancel(_) => ControlMessageType::PublishNamespaceCancel,
      ControlMessage::PublishNamespaceDone(_) => ControlMessageType::PublishNamespaceDone,
      ControlMessage::RequestError(_) => ControlMessageType::RequestError,
      ControlMessage::RequestOk(_) => ControlMessageType::RequestOk,
      ControlMessage::Publish(_) => ControlMessageType::Publish,
      ControlMessage::PublishOk(_) => ControlMessageType::PublishOk,
      ControlMessage::PublishDone(_) => ControlMessageType::PublishDone,
      ControlMessage::ClientSetup(_) => ControlMessageType::ClientSetup,
      ControlMessage::Fetch(_) => ControlMessageType::Fetch,
      ControlMessage::FetchCancel(_) => ControlMessageType::FetchCancel,
      ControlMessage::FetchOk(_) => ControlMessageType::FetchOk,
      ControlMessage::Goaway(_) => ControlMessageType::GoAway,
      ControlMessage::MaxRequestId(_) => ControlMessageType::MaxRequestId,
      ControlMessage::ServerSetup(_) => ControlMessageType::ServerSetup,
      ControlMessage::Subscribe(_) => ControlMessageType::Subscribe,
      ControlMessage::SubscribeOk(_) => ControlMessageType::SubscribeOk,
      ControlMessage::RequestUpdate(_) => ControlMessageType::RequestUpdate,
      ControlMessage::RequestsBlocked(_) => ControlMessageType::RequestsBlocked,
      ControlMessage::TrackStatus(_) => ControlMessageType::TrackStatus,
      ControlMessage::Unsubscribe(_) => ControlMessageType::Unsubscribe,
      ControlMessage::SubscribeNamespace(_) => ControlMessageType::SubscribeNamespace,
      ControlMessage::UnsubscribeNamespace(_) => ControlMessageType::UnsubscribeNamespace,
      ControlMessage::Switch(_) => ControlMessageType::Switch,
    }
  }
}

#[cfg(test)]
mod tests {
  use crate::model::{
    common::tuple::Tuple,
    parameter::{authorization_token::AuthorizationToken, message_parameter::MessageParameter},
  };

  use super::*;

  #[test]
  fn test_announce_roundtrip() {
    let request_id = 12345;
    let track_namespace = Tuple::from_utf8_path("god/dayyum");
    let parameters = vec![MessageParameter::new_authorization_token(
      AuthorizationToken::new_use_value(0, Bytes::from_static(b"test-token")),
    )];
    let announce = PublishNamespace {
      request_id,
      track_namespace,
      parameters,
    };

    let mut buf = announce.serialize().unwrap();
    let deserialized = ControlMessage::deserialize(&mut buf).unwrap();
    if let ControlMessage::PublishNamespace(deserialized_announce) = deserialized {
      assert_eq!(*deserialized_announce, announce);
    } else {
      panic!("Expected ControlMessage::PublishNamespace variant");
    }
    assert!(!buf.has_remaining());
  }
}