1use 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 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 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 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}