Skip to main content

moqtail/model/control/
subscribe.rs

1// Copyright 2025 The MOQtail Authors
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use super::constant::{ControlMessageType, FilterType};
16use super::control_message::ControlMessageTrait;
17use crate::model::common::location::Location;
18use crate::model::common::tuple::{Tuple, TupleField};
19use crate::model::common::varint::{BufMutVarIntExt, BufVarIntExt};
20use crate::model::data::full_track_name::FullTrackName;
21use crate::model::error::ParseError;
22use crate::model::parameter::message_parameter::{
23  MessageParameter, deserialize_message_parameters, serialize_message_parameters,
24};
25use bytes::{Buf, BufMut, Bytes, BytesMut};
26
27#[derive(Debug, PartialEq, Clone)]
28pub struct Subscribe {
29  pub request_id: u64,
30  pub track_namespace: Tuple,
31  pub track_name: TupleField,
32  pub subscribe_parameters: Vec<MessageParameter>,
33}
34
35impl Subscribe {
36  pub fn new(
37    request_id: u64,
38    track_namespace: Tuple,
39    track_name: TupleField,
40    subscribe_parameters: Vec<MessageParameter>,
41  ) -> Self {
42    Self {
43      request_id,
44      track_namespace,
45      track_name,
46      subscribe_parameters,
47    }
48  }
49
50  pub fn new_next_group_start(
51    request_id: u64,
52    track_namespace: Tuple,
53    track_name: TupleField,
54    subscribe_parameters: Vec<MessageParameter>,
55  ) -> Self {
56    let mut params = subscribe_parameters;
57    params.push(MessageParameter::new_subscription_filter(
58      FilterType::NextGroupStart,
59      None,
60      None,
61    ));
62    Self::new(request_id, track_namespace, track_name, params)
63  }
64
65  pub fn new_latest_object(
66    request_id: u64,
67    track_namespace: Tuple,
68    track_name: TupleField,
69    subscribe_parameters: Vec<MessageParameter>,
70  ) -> Self {
71    let mut params = subscribe_parameters;
72    params.push(MessageParameter::new_subscription_filter(
73      FilterType::LatestObject,
74      None,
75      None,
76    ));
77    Self::new(request_id, track_namespace, track_name, params)
78  }
79
80  pub fn new_absolute_start(
81    request_id: u64,
82    track_namespace: Tuple,
83    track_name: TupleField,
84    start_location: Location,
85    subscribe_parameters: Vec<MessageParameter>,
86  ) -> Self {
87    let mut params = subscribe_parameters;
88    params.push(MessageParameter::new_subscription_filter(
89      FilterType::AbsoluteStart,
90      Some(start_location),
91      None,
92    ));
93    Self::new(request_id, track_namespace, track_name, params)
94  }
95
96  pub fn new_absolute_range(
97    request_id: u64,
98    track_namespace: Tuple,
99    track_name: TupleField,
100    start_location: Location,
101    end_group: u64,
102    subscribe_parameters: Vec<MessageParameter>,
103  ) -> Self {
104    assert!(
105      end_group >= start_location.group,
106      "End Group must be >= Start Group"
107    );
108    let mut params = subscribe_parameters;
109    params.push(MessageParameter::new_subscription_filter(
110      FilterType::AbsoluteRange,
111      Some(start_location),
112      Some(end_group),
113    ));
114    Self::new(request_id, track_namespace, track_name, params)
115  }
116
117  pub fn get_full_track_name(&self) -> FullTrackName {
118    FullTrackName {
119      namespace: self.track_namespace.clone(),
120      name: self.track_name.clone(),
121    }
122  }
123
124  /// Returns the SubscriptionFilter parameter if present.
125  pub fn get_subscription_filter(&self) -> Option<(FilterType, Option<Location>, Option<u64>)> {
126    self.subscribe_parameters.iter().find_map(|p| {
127      if let MessageParameter::SubscriptionFilter {
128        filter_type,
129        start_location,
130        end_group,
131      } = p
132      {
133        Some((*filter_type, start_location.clone(), *end_group))
134      } else {
135        None
136      }
137    })
138  }
139}
140
141impl ControlMessageTrait for Subscribe {
142  fn serialize(&self) -> Result<Bytes, ParseError> {
143    let mut buf = BytesMut::new();
144    buf.put_vi(ControlMessageType::Subscribe)?;
145
146    let mut payload = BytesMut::new();
147    payload.put_vi(self.request_id)?;
148
149    payload.extend_from_slice(&self.track_namespace.serialize()?);
150    payload.put_vi(self.track_name.len())?;
151    payload.extend_from_slice(self.track_name.as_bytes());
152
153    payload.put_vi(self.subscribe_parameters.len())?;
154    payload.extend_from_slice(&serialize_message_parameters(&self.subscribe_parameters)?);
155
156    let payload_len: u16 = payload
157      .len()
158      .try_into()
159      .map_err(|e: std::num::TryFromIntError| ParseError::CastingError {
160        context: "Subscribe::serialize",
161        from_type: "usize",
162        to_type: "u16",
163        details: e.to_string(),
164      })?;
165
166    buf.put_u16(payload_len);
167    buf.extend_from_slice(&payload);
168    Ok(buf.freeze())
169  }
170
171  fn parse_payload(payload: &mut Bytes) -> Result<Box<Self>, ParseError> {
172    let request_id = payload.get_vi()?;
173    let track_namespace = Tuple::deserialize(payload)?;
174
175    let name_len_u64 = payload.get_vi()?;
176    let name_len: usize = name_len_u64
177      .try_into()
178      .map_err(|e: std::num::TryFromIntError| ParseError::CastingError {
179        context: "Subscribe::parse_payload(track_name_len)",
180        from_type: "u64",
181        to_type: "usize",
182        details: e.to_string(),
183      })?;
184
185    if payload.remaining() < name_len {
186      return Err(ParseError::NotEnoughBytes {
187        context: "Subscribe::parse_payload(track_name)",
188        needed: name_len,
189        available: payload.remaining(),
190      });
191    }
192    let track_name = TupleField::new(payload.copy_to_bytes(name_len));
193
194    let param_count = payload.get_vi()?;
195    let subscribe_parameters =
196      deserialize_message_parameters(payload, param_count, ControlMessageType::Subscribe)?;
197
198    Ok(Box::new(Subscribe {
199      request_id,
200      track_namespace,
201      track_name,
202      subscribe_parameters,
203    }))
204  }
205
206  fn get_type(&self) -> ControlMessageType {
207    ControlMessageType::Subscribe
208  }
209}
210
211#[cfg(test)]
212mod tests {
213  use super::*;
214  use crate::model::control::constant::GroupOrder;
215  use bytes::Buf;
216
217  #[test]
218  fn test_roundtrip() {
219    let mut subscribe = Subscribe::new_absolute_range(
220      128242,
221      Tuple::from_utf8_path("nein/nein/nein"),
222      TupleField::from_utf8("${Name}"),
223      Location {
224        group: 81,
225        object: 81,
226      },
227      100,
228      vec![
229        MessageParameter::new_subscriber_priority(31),
230        MessageParameter::new_forward(false),
231      ],
232    );
233    // Wire encoding canonicalizes parameter order ascending by type (delta-encoding requirement).
234    subscribe
235      .subscribe_parameters
236      .sort_by_key(|p| p.type_value());
237
238    let mut buf = subscribe.serialize().unwrap();
239    let msg_type = buf.get_vi().unwrap();
240    assert_eq!(msg_type, ControlMessageType::Subscribe as u64);
241    let msg_length = buf.get_u16();
242    assert_eq!(msg_length as usize, buf.remaining());
243    let deserialized = Subscribe::parse_payload(&mut buf).unwrap();
244    assert_eq!(*deserialized, subscribe);
245    assert!(!buf.has_remaining());
246  }
247
248  #[test]
249  fn test_excess_roundtrip() {
250    let mut subscribe = Subscribe::new_absolute_range(
251      128242,
252      Tuple::from_utf8_path("nein/nein/nein"),
253      TupleField::from_utf8("${Name}"),
254      Location {
255        group: 81,
256        object: 81,
257      },
258      100,
259      vec![
260        MessageParameter::new_subscriber_priority(31),
261        MessageParameter::new_group_order(GroupOrder::Ascending),
262        MessageParameter::new_forward(true),
263      ],
264    );
265    // Wire encoding canonicalizes parameter order ascending by type (delta-encoding requirement).
266    subscribe
267      .subscribe_parameters
268      .sort_by_key(|p| p.type_value());
269
270    let serialized = subscribe.serialize().unwrap();
271    let mut excess = BytesMut::new();
272    excess.extend_from_slice(&serialized);
273    excess.extend_from_slice(&[9u8, 1u8, 1u8]);
274    let mut buf = excess.freeze();
275
276    let msg_type = buf.get_vi().unwrap();
277    assert_eq!(msg_type, ControlMessageType::Subscribe as u64);
278    let msg_length = buf.get_u16();
279
280    assert_eq!(msg_length as usize, buf.remaining() - 3);
281    let deserialized = Subscribe::parse_payload(&mut buf).unwrap();
282    assert_eq!(*deserialized, subscribe);
283    assert_eq!(buf.chunk(), &[9u8, 1u8, 1u8]);
284  }
285
286  #[test]
287  fn test_partial_message() {
288    let subscribe = Subscribe::new_absolute_range(
289      128242,
290      Tuple::from_utf8_path("nein/nein/nein"),
291      TupleField::from_utf8("${Name}"),
292      Location {
293        group: 81,
294        object: 81,
295      },
296      100,
297      vec![
298        MessageParameter::new_subscriber_priority(31),
299        MessageParameter::new_group_order(GroupOrder::Ascending),
300        MessageParameter::new_forward(true),
301      ],
302    );
303
304    let mut buf = subscribe.serialize().unwrap();
305    let msg_type = buf.get_vi().unwrap();
306    assert_eq!(msg_type, ControlMessageType::Subscribe as u64);
307    let msg_length = buf.get_u16();
308    assert_eq!(msg_length as usize, buf.remaining());
309
310    let upper = buf.remaining() / 2;
311    let mut partial = buf.slice(..upper);
312    let deserialized = Subscribe::parse_payload(&mut partial);
313    assert!(deserialized.is_err());
314  }
315}