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 for_each_offered_endpoint<F>(&self, mut f: F)
179 where
180 F: FnMut(crate::OfferedEndpoint),
181 {
182 let header = match &self.kind {
183 RawPayloadKind::Sd(header) => header,
184 RawPayloadKind::Raw(_) => return,
185 };
186 for entry in &header.entries {
187 if let sd::Entry::OfferService(svc) | sd::Entry::StopOfferService(svc) = entry {
188 let is_offer = matches!(entry, sd::Entry::OfferService(_));
189 let endpoint =
190 sd::extract_ipv4_endpoint(&header.options).map(|(addr, protocol)| {
191 crate::NetEndpoint::new(core::net::SocketAddr::V4(addr), protocol)
192 });
193 f(crate::OfferedEndpoint {
194 service_id: svc.service_id,
195 instance_id: svc.instance_id,
196 major_version: svc.major_version,
197 minor_version: svc.minor_version,
198 endpoint,
199 is_offer,
200 });
201 }
202 }
203 }
204
205 fn for_each_service_instance<F>(&self, mut f: F)
206 where
207 F: FnMut(u16, u16),
208 {
209 let header = match &self.kind {
210 RawPayloadKind::Sd(header) => header,
211 RawPayloadKind::Raw(_) => return,
212 };
213 for entry in &header.entries {
214 let (svc, inst) = match entry {
215 sd::Entry::FindService(svc)
216 | sd::Entry::OfferService(svc)
217 | sd::Entry::StopOfferService(svc) => (svc.service_id, svc.instance_id),
218 sd::Entry::SubscribeEventGroup(eg) | sd::Entry::SubscribeAckEventGroup(eg) => {
219 (eg.service_id, eg.instance_id)
220 }
221 };
222 f(svc, inst);
223 }
224 }
225}
226
227#[cfg(test)]
228mod tests {
229 use super::*;
230 use crate::traits::WireFormat;
231 use std::net::Ipv4Addr;
232
233 fn make_sd_payload() -> RawPayload {
234 let header = VecSdHeader {
235 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
236 entries: std::vec![],
237 options: std::vec![],
238 };
239 RawPayload::new_sd_payload(&header)
240 }
241
242 fn make_raw_payload() -> RawPayload {
243 let mid = MessageId::new_from_service_and_method(0x1234, 0x0001);
244 RawPayload::from_payload_bytes(mid, &[0xDE, 0xAD]).unwrap()
245 }
246
247 #[test]
248 fn set_reboot_flag_flips_reboot_and_preserves_unicast() {
249 let mut header = VecSdHeader {
251 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
252 entries: std::vec![],
253 options: std::vec![],
254 };
255 assert_eq!(header.flags.reboot(), sd::RebootFlag::RecentlyRebooted);
256 assert!(header.flags.unicast());
257
258 RawPayload::set_reboot_flag(&mut header, sd::RebootFlag::Continuous);
259 assert_eq!(header.flags.reboot(), sd::RebootFlag::Continuous);
260 assert!(
261 header.flags.unicast(),
262 "unicast bit must not be disturbed by set_reboot_flag"
263 );
264
265 RawPayload::set_reboot_flag(&mut header, sd::RebootFlag::RecentlyRebooted);
266 assert_eq!(header.flags.reboot(), sd::RebootFlag::RecentlyRebooted);
267 assert!(header.flags.unicast());
268 }
269
270 #[test]
271 fn set_reboot_flag_preserves_cleared_unicast() {
272 let mut header = VecSdHeader {
275 flags: sd::Flags::new(true, false),
276 entries: std::vec![],
277 options: std::vec![],
278 };
279 RawPayload::set_reboot_flag(&mut header, sd::RebootFlag::Continuous);
280 assert_eq!(header.flags.reboot(), sd::RebootFlag::Continuous);
281 assert!(!header.flags.unicast());
282 }
283
284 #[test]
285 fn raw_bytes_returns_some_for_raw_payload() {
286 let p = make_raw_payload();
287 assert_eq!(p.raw_bytes(), Some(&[0xDE, 0xAD][..]));
288 }
289
290 #[test]
291 fn raw_bytes_returns_none_for_sd_payload() {
292 let p = make_sd_payload();
293 assert_eq!(p.raw_bytes(), None);
294 }
295
296 #[test]
297 fn as_sd_header_returns_some_for_sd() {
298 let p = make_sd_payload();
299 assert!(p.as_sd_header().is_some());
300 }
301
302 #[test]
303 fn as_sd_header_returns_none_for_raw() {
304 let p = make_raw_payload();
305 assert!(p.as_sd_header().is_none());
306 }
307
308 #[test]
309 fn sd_flags_returns_some_for_sd() {
310 let p = make_sd_payload();
311 let flags = p.sd_flags().unwrap();
312 assert!(flags.unicast());
313 }
314
315 #[test]
316 fn sd_flags_returns_none_for_raw() {
317 let p = make_raw_payload();
318 assert!(p.sd_flags().is_none());
319 }
320
321 #[test]
322 fn message_id_correct() {
323 let p = make_raw_payload();
324 assert_eq!(p.message_id().service_id(), 0x1234);
325
326 let sd = make_sd_payload();
327 assert_eq!(sd.message_id(), MessageId::SD);
328 }
329
330 #[test]
331 fn required_size_raw() {
332 let p = make_raw_payload();
333 assert_eq!(p.required_size(), 2);
334 }
335
336 #[test]
337 fn encode_raw_payload() {
338 let p = make_raw_payload();
339 let mut buf = std::vec![0u8; p.required_size()];
340 let n = p.encode(&mut buf.as_mut_slice()).unwrap();
341 assert_eq!(n, 2);
342 assert_eq!(&buf, &[0xDE, 0xAD]);
343 }
344
345 #[test]
346 fn encode_sd_payload() {
347 let p = make_sd_payload();
348 let mut buf = std::vec![0u8; p.required_size()];
349 let n = p.encode(&mut buf.as_mut_slice()).unwrap();
350 assert_eq!(n, p.required_size());
351 }
352
353 #[test]
354 fn from_payload_bytes_sd_roundtrip() {
355 let entry = sd::Entry::FindService(sd::ServiceEntry::find(0x5B));
357 let entries = [entry];
358 let header = sd::Header::new(
359 sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
360 &entries,
361 &[],
362 );
363 let mut buf = std::vec![0u8; header.required_size()];
364 header.encode(&mut buf.as_mut_slice()).unwrap();
365
366 let p = RawPayload::from_payload_bytes(MessageId::SD, &buf).unwrap();
367 assert!(p.as_sd_header().is_some());
368 let sd = p.as_sd_header().unwrap();
369 assert_eq!(sd.entries.len(), 1);
370 }
371
372 #[test]
373 fn from_payload_bytes_non_sd() {
374 let mid = MessageId::new_from_service_and_method(0x5B, 0x01);
375 let p = RawPayload::from_payload_bytes(mid, &[1, 2, 3]).unwrap();
376 assert_eq!(p.raw_bytes(), Some(&[1, 2, 3][..]));
377 }
378
379 #[test]
380 fn new_subscription_sd_header_structure() {
381 let header = RawPayload::new_subscription_sd_header(
382 0x5B,
383 1,
384 1,
385 3,
386 0x01,
387 Ipv4Addr::LOCALHOST,
388 sd::TransportProtocol::Udp,
389 12345,
390 sd::RebootFlag::Continuous,
391 );
392 assert_eq!(header.entries.len(), 1);
393 assert_eq!(header.options.len(), 1);
394 assert!(header.flags.unicast());
395 assert_eq!(header.flags.reboot(), sd::RebootFlag::Continuous);
396
397 let header_reboot = RawPayload::new_subscription_sd_header(
398 0x5B,
399 1,
400 1,
401 3,
402 0x01,
403 Ipv4Addr::LOCALHOST,
404 sd::TransportProtocol::Udp,
405 12345,
406 sd::RebootFlag::RecentlyRebooted,
407 );
408 assert_eq!(
409 header_reboot.flags.reboot(),
410 sd::RebootFlag::RecentlyRebooted
411 );
412 }
413
414 #[test]
415 fn offered_endpoints_from_raw_returns_empty() {
416 let p = make_raw_payload();
417 assert!(p.offered_endpoints().is_empty());
418 }
419
420 fn make_offer_entry(service_id: u16, instance_id: u16) -> sd::ServiceEntry {
421 sd::ServiceEntry {
422 index_first_options_run: 0,
423 index_second_options_run: 0,
424 options_count: sd::OptionsCount::new(1, 0),
425 service_id,
426 instance_id,
427 major_version: 1,
428 ttl: 100,
429 minor_version: 0,
430 }
431 }
432
433 #[test]
434 fn offered_endpoints_with_offer_service() {
435 let offer = sd::Entry::OfferService(make_offer_entry(0x5B, 1));
436 let endpoint = sd::Options::IpV4Endpoint {
437 ip: Ipv4Addr::LOCALHOST,
438 protocol: sd::TransportProtocol::Udp,
439 port: 30000,
440 };
441 let header = VecSdHeader {
442 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
443 entries: std::vec![offer],
444 options: std::vec![endpoint],
445 };
446 let p = RawPayload::new_sd_payload(&header);
447 let endpoints = p.offered_endpoints();
448 assert_eq!(endpoints.len(), 1);
449 assert_eq!(endpoints[0].service_id, 0x5B);
450 assert!(endpoints[0].is_offer);
451 let ep = endpoints[0].endpoint.expect("endpoint option present");
452 assert_eq!(ep.protocol, crate::TransportProtocol::Udp);
453 assert_eq!(
454 ep.addr,
455 core::net::SocketAddr::V4(core::net::SocketAddrV4::new(Ipv4Addr::LOCALHOST, 30000))
456 );
457 }
458
459 #[test]
460 fn offered_endpoints_with_stop_offer() {
461 let mut entry = make_offer_entry(0x5B, 1);
462 entry.ttl = 0;
463 let stop = sd::Entry::StopOfferService(entry);
464 let header = VecSdHeader {
465 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
466 entries: std::vec![stop],
467 options: std::vec![],
468 };
469 let p = RawPayload::new_sd_payload(&header);
470 let endpoints = p.offered_endpoints();
471 assert_eq!(endpoints.len(), 1);
472 assert!(!endpoints[0].is_offer);
473 assert!(endpoints[0].endpoint.is_none());
474 }
475
476 #[test]
477 fn offered_endpoints_ignores_non_offer_entries() {
478 let find = sd::Entry::FindService(sd::ServiceEntry::find(0x5B));
479 let header = VecSdHeader {
480 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
481 entries: std::vec![find],
482 options: std::vec![],
483 };
484 let p = RawPayload::new_sd_payload(&header);
485 assert!(p.offered_endpoints().is_empty());
486 }
487
488 #[test]
489 fn service_instances_returns_all_entry_types() {
490 let offer = sd::Entry::OfferService(make_offer_entry(0x47, 1));
491 let find = sd::Entry::FindService(sd::ServiceEntry::find(0x5D));
492 let sub = sd::Entry::SubscribeEventGroup(sd::EventGroupEntry::new(0x5B, 1, 1, 3, 0x01));
493 let header = VecSdHeader {
494 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
495 entries: std::vec![offer, find, sub],
496 options: std::vec![],
497 };
498 let p = RawPayload::new_sd_payload(&header);
499 let instances = p.service_instances();
500 assert_eq!(instances.len(), 3);
501 assert_eq!(instances[0], (0x47, 1)); assert_eq!(instances[1], (0x5D, 0xFFFF)); assert_eq!(instances[2], (0x5B, 1)); }
505
506 #[test]
507 fn service_instances_empty_for_raw_payload() {
508 let mid = MessageId::new_from_service_and_method(0x5B, 0x01);
509 let p = RawPayload::from_payload_bytes(mid, &[1, 2, 3]).unwrap();
510 assert!(p.service_instances().is_empty());
511 }
512
513 #[test]
514 fn service_instances_empty_for_no_entries() {
515 let header = VecSdHeader {
516 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
517 entries: std::vec![],
518 options: std::vec![],
519 };
520 let p = RawPayload::new_sd_payload(&header);
521 assert!(p.service_instances().is_empty());
522 }
523
524 #[test]
525 fn vec_sd_header_required_size_and_encode() {
526 let header = VecSdHeader {
527 flags: sd::Flags::new_sd(sd::RebootFlag::RecentlyRebooted),
528 entries: std::vec![],
529 options: std::vec![],
530 };
531 let size = header.required_size();
532 assert!(size > 0);
533 let mut buf = std::vec![0u8; size];
534 let n = header.encode(&mut buf.as_mut_slice()).unwrap();
535 assert_eq!(n, size);
536 }
537}