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