1use std::vec::Vec;
11
12use embedded_io::Error as _;
13
14use crate::protocol::{self, MessageId, sd};
15use crate::traits::{PayloadWireFormat, WireFormat};
16
17#[derive(Clone, Debug, Eq, PartialEq)]
19pub struct VecSdHeader {
20 pub flags: sd::Flags,
22 pub entries: Vec<sd::Entry>,
24 pub options: Vec<sd::Options>,
26}
27
28impl WireFormat for VecSdHeader {
29 fn required_size(&self) -> usize {
30 sd::Header::new(self.flags, &self.entries, &self.options).required_size()
31 }
32
33 fn encode<T: embedded_io::Write>(&self, writer: &mut T) -> Result<usize, protocol::Error> {
34 sd::Header::new(self.flags, &self.entries, &self.options).encode(writer)
35 }
36}
37
38#[derive(Clone, Debug, Eq, PartialEq)]
40enum RawPayloadKind {
41 Sd(VecSdHeader),
43 Raw(Vec<u8>),
45}
46
47#[derive(Clone, Debug, Eq, PartialEq)]
54pub struct RawPayload {
55 message_id: MessageId,
56 kind: RawPayloadKind,
57}
58
59impl RawPayload {
60 #[must_use]
62 pub fn raw_bytes(&self) -> Option<&[u8]> {
63 match &self.kind {
64 RawPayloadKind::Raw(bytes) => Some(bytes),
65 RawPayloadKind::Sd(_) => None,
66 }
67 }
68}
69
70impl PayloadWireFormat for RawPayload {
71 type SdHeader = VecSdHeader;
72
73 fn message_id(&self) -> MessageId {
74 self.message_id
75 }
76
77 fn as_sd_header(&self) -> Option<&VecSdHeader> {
78 match &self.kind {
79 RawPayloadKind::Sd(header) => Some(header),
80 RawPayloadKind::Raw(_) => None,
81 }
82 }
83
84 fn from_payload_bytes(message_id: MessageId, payload: &[u8]) -> Result<Self, protocol::Error> {
85 if message_id == MessageId::SD {
86 let view = sd::SdHeaderView::parse(payload)?;
87 let mut entries = Vec::new();
88 for ev in view.entries() {
89 entries.push(ev.to_owned()?);
90 }
91 let mut options = Vec::new();
92 for ov in view.options() {
93 options.push(ov.to_owned()?);
94 }
95 Ok(Self {
96 message_id,
97 kind: RawPayloadKind::Sd(VecSdHeader {
98 flags: view.flags(),
99 entries,
100 options,
101 }),
102 })
103 } else {
104 Ok(Self {
105 message_id,
106 kind: RawPayloadKind::Raw(payload.to_vec()),
107 })
108 }
109 }
110
111 fn new_sd_payload(header: &VecSdHeader) -> Self {
112 Self {
113 message_id: MessageId::SD,
114 kind: RawPayloadKind::Sd(header.clone()),
115 }
116 }
117
118 fn sd_flags(&self) -> Option<sd::Flags> {
119 match &self.kind {
120 RawPayloadKind::Sd(header) => Some(header.flags),
121 RawPayloadKind::Raw(_) => None,
122 }
123 }
124
125 fn required_size(&self) -> usize {
126 match &self.kind {
127 RawPayloadKind::Sd(header) => header.required_size(),
128 RawPayloadKind::Raw(bytes) => bytes.len(),
129 }
130 }
131
132 fn encode<T: embedded_io::Write>(&self, writer: &mut T) -> Result<usize, protocol::Error> {
133 match &self.kind {
134 RawPayloadKind::Sd(header) => header.encode(writer),
135 RawPayloadKind::Raw(bytes) => {
136 writer
137 .write_all(bytes)
138 .map_err(|e| protocol::Error::Io(e.kind()))?;
139 Ok(bytes.len())
140 }
141 }
142 }
143
144 fn new_subscription_sd_header(
145 service_id: u16,
146 instance_id: u16,
147 major_version: u8,
148 ttl: u32,
149 event_group_id: u16,
150 client_ip: std::net::Ipv4Addr,
151 protocol: sd::TransportProtocol,
152 client_port: u16,
153 reboot_flag: sd::RebootFlag,
154 ) -> VecSdHeader {
155 let entry = sd::Entry::SubscribeEventGroup(sd::EventGroupEntry::new(
156 service_id,
157 instance_id,
158 major_version,
159 ttl,
160 event_group_id,
161 ));
162 let endpoint = sd::Options::IpV4Endpoint {
163 ip: client_ip,
164 protocol,
165 port: client_port,
166 };
167 VecSdHeader {
168 flags: sd::Flags::new_sd(reboot_flag),
169 entries: std::vec![entry],
170 options: std::vec![endpoint],
171 }
172 }
173
174 fn set_reboot_flag(header: &mut VecSdHeader, reboot: sd::RebootFlag) {
175 header.flags = sd::Flags::new(bool::from(reboot), header.flags.unicast());
176 }
177
178 fn offered_endpoints(&self) -> Vec<crate::OfferedEndpoint> {
179 let header = match &self.kind {
180 RawPayloadKind::Sd(header) => header,
181 RawPayloadKind::Raw(_) => return Vec::new(),
182 };
183 header
184 .entries
185 .iter()
186 .filter_map(|entry| match entry {
187 sd::Entry::OfferService(svc) | sd::Entry::StopOfferService(svc) => {
188 let is_offer = matches!(entry, sd::Entry::OfferService(_));
189 let addr = sd::extract_ipv4_endpoint(&header.options);
190 Some(crate::OfferedEndpoint {
191 service_id: svc.service_id,
192 instance_id: svc.instance_id,
193 major_version: svc.major_version,
194 minor_version: svc.minor_version,
195 addr,
196 is_offer,
197 })
198 }
199 _ => None,
200 })
201 .collect()
202 }
203
204 fn service_instances(&self) -> Vec<(u16, u16)> {
205 let header = match &self.kind {
206 RawPayloadKind::Sd(header) => header,
207 RawPayloadKind::Raw(_) => return Vec::new(),
208 };
209 header
210 .entries
211 .iter()
212 .map(|entry| match entry {
213 sd::Entry::FindService(svc)
214 | sd::Entry::OfferService(svc)
215 | sd::Entry::StopOfferService(svc) => (svc.service_id, svc.instance_id),
216 sd::Entry::SubscribeEventGroup(eg) | sd::Entry::SubscribeAckEventGroup(eg) => {
217 (eg.service_id, eg.instance_id)
218 }
219 })
220 .collect()
221 }
222}
223
224#[cfg(test)]
225mod tests {
226 use super::*;
227 use crate::traits::WireFormat;
228 use std::net::Ipv4Addr;
229
230 fn make_sd_payload() -> RawPayload {
231 let header = VecSdHeader {
232 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
233 entries: std::vec![],
234 options: std::vec![],
235 };
236 RawPayload::new_sd_payload(&header)
237 }
238
239 fn make_raw_payload() -> RawPayload {
240 let mid = MessageId::new_from_service_and_method(0x1234, 0x0001);
241 RawPayload::from_payload_bytes(mid, &[0xDE, 0xAD]).unwrap()
242 }
243
244 #[test]
245 fn set_reboot_flag_flips_reboot_and_preserves_unicast() {
246 let mut header = VecSdHeader {
248 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
249 entries: std::vec![],
250 options: std::vec![],
251 };
252 assert_eq!(header.flags.reboot(), sd::RebootFlag::RecentlyRebooted);
253 assert!(header.flags.unicast());
254
255 RawPayload::set_reboot_flag(&mut header, sd::RebootFlag::Continuous);
256 assert_eq!(header.flags.reboot(), sd::RebootFlag::Continuous);
257 assert!(
258 header.flags.unicast(),
259 "unicast bit must not be disturbed by set_reboot_flag"
260 );
261
262 RawPayload::set_reboot_flag(&mut header, sd::RebootFlag::RecentlyRebooted);
263 assert_eq!(header.flags.reboot(), sd::RebootFlag::RecentlyRebooted);
264 assert!(header.flags.unicast());
265 }
266
267 #[test]
268 fn set_reboot_flag_preserves_cleared_unicast() {
269 let mut header = VecSdHeader {
272 flags: sd::Flags::new(true, false),
273 entries: std::vec![],
274 options: std::vec![],
275 };
276 RawPayload::set_reboot_flag(&mut header, sd::RebootFlag::Continuous);
277 assert_eq!(header.flags.reboot(), sd::RebootFlag::Continuous);
278 assert!(!header.flags.unicast());
279 }
280
281 #[test]
282 fn raw_bytes_returns_some_for_raw_payload() {
283 let p = make_raw_payload();
284 assert_eq!(p.raw_bytes(), Some(&[0xDE, 0xAD][..]));
285 }
286
287 #[test]
288 fn raw_bytes_returns_none_for_sd_payload() {
289 let p = make_sd_payload();
290 assert_eq!(p.raw_bytes(), None);
291 }
292
293 #[test]
294 fn as_sd_header_returns_some_for_sd() {
295 let p = make_sd_payload();
296 assert!(p.as_sd_header().is_some());
297 }
298
299 #[test]
300 fn as_sd_header_returns_none_for_raw() {
301 let p = make_raw_payload();
302 assert!(p.as_sd_header().is_none());
303 }
304
305 #[test]
306 fn sd_flags_returns_some_for_sd() {
307 let p = make_sd_payload();
308 let flags = p.sd_flags().unwrap();
309 assert!(flags.unicast());
310 }
311
312 #[test]
313 fn sd_flags_returns_none_for_raw() {
314 let p = make_raw_payload();
315 assert!(p.sd_flags().is_none());
316 }
317
318 #[test]
319 fn message_id_correct() {
320 let p = make_raw_payload();
321 assert_eq!(p.message_id().service_id(), 0x1234);
322
323 let sd = make_sd_payload();
324 assert_eq!(sd.message_id(), MessageId::SD);
325 }
326
327 #[test]
328 fn required_size_raw() {
329 let p = make_raw_payload();
330 assert_eq!(p.required_size(), 2);
331 }
332
333 #[test]
334 fn encode_raw_payload() {
335 let p = make_raw_payload();
336 let mut buf = std::vec![0u8; p.required_size()];
337 let n = p.encode(&mut buf.as_mut_slice()).unwrap();
338 assert_eq!(n, 2);
339 assert_eq!(&buf, &[0xDE, 0xAD]);
340 }
341
342 #[test]
343 fn encode_sd_payload() {
344 let p = make_sd_payload();
345 let mut buf = std::vec![0u8; p.required_size()];
346 let n = p.encode(&mut buf.as_mut_slice()).unwrap();
347 assert_eq!(n, p.required_size());
348 }
349
350 #[test]
351 fn from_payload_bytes_sd_roundtrip() {
352 let entry = sd::Entry::FindService(sd::ServiceEntry::find(0x5B));
354 let entries = [entry];
355 let header = sd::Header::new(
356 sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
357 &entries,
358 &[],
359 );
360 let mut buf = std::vec![0u8; header.required_size()];
361 header.encode(&mut buf.as_mut_slice()).unwrap();
362
363 let p = RawPayload::from_payload_bytes(MessageId::SD, &buf).unwrap();
364 assert!(p.as_sd_header().is_some());
365 let sd = p.as_sd_header().unwrap();
366 assert_eq!(sd.entries.len(), 1);
367 }
368
369 #[test]
370 fn from_payload_bytes_non_sd() {
371 let mid = MessageId::new_from_service_and_method(0x5B, 0x01);
372 let p = RawPayload::from_payload_bytes(mid, &[1, 2, 3]).unwrap();
373 assert_eq!(p.raw_bytes(), Some(&[1, 2, 3][..]));
374 }
375
376 #[test]
377 fn new_subscription_sd_header_structure() {
378 let header = RawPayload::new_subscription_sd_header(
379 0x5B,
380 1,
381 1,
382 3,
383 0x01,
384 Ipv4Addr::LOCALHOST,
385 sd::TransportProtocol::Udp,
386 12345,
387 sd::RebootFlag::Continuous,
388 );
389 assert_eq!(header.entries.len(), 1);
390 assert_eq!(header.options.len(), 1);
391 assert!(header.flags.unicast());
392 assert_eq!(header.flags.reboot(), sd::RebootFlag::Continuous);
393
394 let header_reboot = RawPayload::new_subscription_sd_header(
395 0x5B,
396 1,
397 1,
398 3,
399 0x01,
400 Ipv4Addr::LOCALHOST,
401 sd::TransportProtocol::Udp,
402 12345,
403 sd::RebootFlag::RecentlyRebooted,
404 );
405 assert_eq!(
406 header_reboot.flags.reboot(),
407 sd::RebootFlag::RecentlyRebooted
408 );
409 }
410
411 #[test]
412 fn offered_endpoints_from_raw_returns_empty() {
413 let p = make_raw_payload();
414 assert!(p.offered_endpoints().is_empty());
415 }
416
417 fn make_offer_entry(service_id: u16, instance_id: u16) -> sd::ServiceEntry {
418 sd::ServiceEntry {
419 index_first_options_run: 0,
420 index_second_options_run: 0,
421 options_count: sd::OptionsCount::new(1, 0),
422 service_id,
423 instance_id,
424 major_version: 1,
425 ttl: 100,
426 minor_version: 0,
427 }
428 }
429
430 #[test]
431 fn offered_endpoints_with_offer_service() {
432 let offer = sd::Entry::OfferService(make_offer_entry(0x5B, 1));
433 let endpoint = sd::Options::IpV4Endpoint {
434 ip: Ipv4Addr::LOCALHOST,
435 protocol: sd::TransportProtocol::Udp,
436 port: 30000,
437 };
438 let header = VecSdHeader {
439 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
440 entries: std::vec![offer],
441 options: std::vec![endpoint],
442 };
443 let p = RawPayload::new_sd_payload(&header);
444 let endpoints = p.offered_endpoints();
445 assert_eq!(endpoints.len(), 1);
446 assert_eq!(endpoints[0].service_id, 0x5B);
447 assert!(endpoints[0].is_offer);
448 assert!(endpoints[0].addr.is_some());
449 }
450
451 #[test]
452 fn offered_endpoints_with_stop_offer() {
453 let mut entry = make_offer_entry(0x5B, 1);
454 entry.ttl = 0;
455 let stop = sd::Entry::StopOfferService(entry);
456 let header = VecSdHeader {
457 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
458 entries: std::vec![stop],
459 options: std::vec![],
460 };
461 let p = RawPayload::new_sd_payload(&header);
462 let endpoints = p.offered_endpoints();
463 assert_eq!(endpoints.len(), 1);
464 assert!(!endpoints[0].is_offer);
465 assert!(endpoints[0].addr.is_none());
466 }
467
468 #[test]
469 fn offered_endpoints_ignores_non_offer_entries() {
470 let find = sd::Entry::FindService(sd::ServiceEntry::find(0x5B));
471 let header = VecSdHeader {
472 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
473 entries: std::vec![find],
474 options: std::vec![],
475 };
476 let p = RawPayload::new_sd_payload(&header);
477 assert!(p.offered_endpoints().is_empty());
478 }
479
480 #[test]
481 fn service_instances_returns_all_entry_types() {
482 let offer = sd::Entry::OfferService(make_offer_entry(0x47, 1));
483 let find = sd::Entry::FindService(sd::ServiceEntry::find(0x5D));
484 let sub = sd::Entry::SubscribeEventGroup(sd::EventGroupEntry::new(0x5B, 1, 1, 3, 0x01));
485 let header = VecSdHeader {
486 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
487 entries: std::vec![offer, find, sub],
488 options: std::vec![],
489 };
490 let p = RawPayload::new_sd_payload(&header);
491 let instances = p.service_instances();
492 assert_eq!(instances.len(), 3);
493 assert_eq!(instances[0], (0x47, 1)); assert_eq!(instances[1], (0x5D, 0xFFFF)); assert_eq!(instances[2], (0x5B, 1)); }
497
498 #[test]
499 fn service_instances_empty_for_raw_payload() {
500 let mid = MessageId::new_from_service_and_method(0x5B, 0x01);
501 let p = RawPayload::from_payload_bytes(mid, &[1, 2, 3]).unwrap();
502 assert!(p.service_instances().is_empty());
503 }
504
505 #[test]
506 fn service_instances_empty_for_no_entries() {
507 let header = VecSdHeader {
508 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
509 entries: std::vec![],
510 options: std::vec![],
511 };
512 let p = RawPayload::new_sd_payload(&header);
513 assert!(p.service_instances().is_empty());
514 }
515
516 #[test]
517 fn vec_sd_header_required_size_and_encode() {
518 let header = VecSdHeader {
519 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
520 entries: std::vec![],
521 options: std::vec![],
522 };
523 let size = header.required_size();
524 assert!(size > 0);
525 let mut buf = std::vec![0u8; size];
526 let n = header.encode(&mut buf.as_mut_slice()).unwrap();
527 assert_eq!(n, size);
528 }
529}