1use crate::{
2 protocol::{Error, MessageId, MessageTypeField, ReturnCode, byte_order::WriteBytesExt},
3 traits::WireFormat,
4};
5
6#[derive(Clone, Debug, Eq, PartialEq)]
8pub struct Header {
9 message_id: MessageId,
11 length: u32,
14 request_id: u32,
16 protocol_version: u8,
17 interface_version: u8,
18 message_type: MessageTypeField,
19 return_code: ReturnCode,
20}
21
22impl Header {
23 #[must_use]
25 pub const fn message_id(&self) -> MessageId {
26 self.message_id
27 }
28
29 #[must_use]
31 pub const fn length(&self) -> u32 {
32 self.length
33 }
34
35 #[must_use]
37 pub const fn request_id(&self) -> u32 {
38 self.request_id
39 }
40
41 #[must_use]
43 pub const fn protocol_version(&self) -> u8 {
44 self.protocol_version
45 }
46
47 #[must_use]
49 pub const fn interface_version(&self) -> u8 {
50 self.interface_version
51 }
52
53 #[must_use]
55 pub const fn message_type(&self) -> MessageTypeField {
56 self.message_type
57 }
58
59 #[must_use]
61 pub const fn return_code(&self) -> ReturnCode {
62 self.return_code
63 }
64
65 #[must_use]
73 pub const fn upper_header_bytes(&self) -> [u8; 8] {
74 let rid = self.request_id.to_be_bytes();
75 [
76 rid[0],
77 rid[1],
78 rid[2],
79 rid[3],
80 self.protocol_version,
81 self.interface_version,
82 self.message_type.as_u8(),
83 self.return_code.as_u8(),
84 ]
85 }
86
87 #[must_use]
94 pub const fn from_fields(
95 message_id: MessageId,
96 length: u32,
97 request_id: u32,
98 protocol_version: u8,
99 interface_version: u8,
100 message_type: MessageTypeField,
101 return_code: ReturnCode,
102 ) -> Self {
103 Self {
104 message_id,
105 length,
106 request_id,
107 protocol_version,
108 interface_version,
109 message_type,
110 return_code,
111 }
112 }
113
114 #[must_use]
120 #[allow(clippy::cast_possible_truncation)]
121 pub const fn new(
122 message_id: MessageId,
123 request_id: u32,
124 protocol_version: u8,
125 interface_version: u8,
126 message_type: MessageTypeField,
127 return_code: ReturnCode,
128 payload_len: usize,
129 ) -> Self {
130 assert!(payload_len <= u32::MAX as usize - 8);
131 Self {
132 message_id,
133 length: 8 + payload_len as u32,
134 request_id,
135 protocol_version,
136 interface_version,
137 message_type,
138 return_code,
139 }
140 }
141
142 #[must_use]
148 #[allow(clippy::cast_possible_truncation)]
149 pub const fn new_sd(request_id: u32, sd_header_size: usize) -> Self {
150 assert!(sd_header_size <= u32::MAX as usize - 8);
151 Self {
152 message_id: MessageId::SD,
153 length: 8 + sd_header_size as u32,
154 request_id,
155 protocol_version: 0x01,
156 interface_version: 0x01,
157 message_type: MessageTypeField::new_sd(),
158 return_code: ReturnCode::Ok,
159 }
160 }
161
162 #[must_use]
168 #[allow(clippy::cast_possible_truncation)]
169 pub const fn new_event(
170 service_id: u16,
171 event_id: u16,
172 request_id: u32,
173 protocol_version: u8,
174 interface_version: u8,
175 payload_len: usize,
176 ) -> Self {
177 assert!(payload_len <= u32::MAX as usize - 8);
178 Self {
179 message_id: MessageId::new_from_service_and_method(service_id, event_id),
180 length: 8 + payload_len as u32,
181 request_id,
182 protocol_version,
183 interface_version,
184 message_type: MessageTypeField::new(crate::protocol::MessageType::Notification, false),
185 return_code: ReturnCode::Ok,
186 }
187 }
188
189 #[must_use]
191 pub const fn is_sd(&self) -> bool {
192 self.message_id.is_sd()
193 }
194
195 #[must_use]
197 pub const fn payload_size(&self) -> usize {
198 self.length as usize - 8
199 }
200
201 pub const fn set_request_id(&mut self, request_id: u32) {
203 self.request_id = request_id;
204 }
205}
206
207#[derive(Clone, Copy, Debug)]
209pub struct HeaderView<'a>(&'a [u8; 16]);
210
211impl<'a> HeaderView<'a> {
212 pub fn parse(buf: &'a [u8]) -> Result<(Self, &'a [u8]), Error> {
224 if buf.len() < 16 {
225 return Err(Error::UnexpectedEof);
226 }
227 let header_bytes: &[u8; 16] = buf[..16].try_into().expect("length checked above");
228 let view = Self(header_bytes);
229
230 let pv = view.protocol_version();
232 if pv != 0x01 {
233 return Err(Error::InvalidProtocolVersion(pv));
234 }
235 MessageTypeField::try_from(header_bytes[14])?;
237 ReturnCode::try_from(header_bytes[15])?;
239
240 Ok((view, &buf[16..]))
241 }
242
243 #[must_use]
245 pub fn message_id(&self) -> MessageId {
246 MessageId::from(u32::from_be_bytes([
247 self.0[0], self.0[1], self.0[2], self.0[3],
248 ]))
249 }
250
251 #[must_use]
253 pub fn length(&self) -> u32 {
254 u32::from_be_bytes([self.0[4], self.0[5], self.0[6], self.0[7]])
255 }
256
257 #[must_use]
259 pub fn request_id(&self) -> u32 {
260 u32::from_be_bytes([self.0[8], self.0[9], self.0[10], self.0[11]])
261 }
262
263 #[must_use]
265 pub fn payload_size(&self) -> usize {
266 self.length() as usize - 8
267 }
268
269 #[must_use]
273 pub const fn upper_header_bytes(&self) -> [u8; 8] {
274 [
275 self.0[8], self.0[9], self.0[10], self.0[11], self.0[12], self.0[13], self.0[14],
276 self.0[15],
277 ]
278 }
279
280 #[must_use]
282 pub fn protocol_version(&self) -> u8 {
283 self.0[12]
284 }
285
286 #[must_use]
288 pub fn interface_version(&self) -> u8 {
289 self.0[13]
290 }
291
292 #[must_use]
298 pub fn message_type(&self) -> MessageTypeField {
299 MessageTypeField::try_from(self.0[14]).expect("validated in parse")
301 }
302
303 #[must_use]
309 pub fn return_code(&self) -> ReturnCode {
310 ReturnCode::try_from(self.0[15]).expect("validated in parse")
312 }
313
314 #[must_use]
316 pub fn is_sd(&self) -> bool {
317 self.message_id().is_sd()
318 }
319
320 #[must_use]
322 pub fn to_owned(&self) -> Header {
323 Header {
324 message_id: self.message_id(),
325 length: self.length(),
326 request_id: self.request_id(),
327 protocol_version: self.protocol_version(),
328 interface_version: self.interface_version(),
329 message_type: self.message_type(),
330 return_code: self.return_code(),
331 }
332 }
333}
334
335impl WireFormat for Header {
336 fn required_size(&self) -> usize {
337 16
338 }
339
340 fn encode<T: embedded_io::Write>(&self, writer: &mut T) -> Result<usize, Error> {
341 writer.write_u32_be(self.message_id.message_id())?;
342 writer.write_u32_be(self.length)?;
343 writer.write_u32_be(self.request_id)?;
344 writer.write_u8(self.protocol_version)?;
345 writer.write_u8(self.interface_version)?;
346 writer.write_u8(u8::from(self.message_type))?;
347 writer.write_u8(u8::from(self.return_code))?;
348 Ok(16)
349 }
350}
351
352#[cfg(test)]
353mod tests {
354 use super::*;
355 use crate::protocol::{Error, MessageId, MessageTypeField, ReturnCode};
356
357 fn make_header() -> Header {
358 Header {
359 message_id: MessageId::new_from_service_and_method(0x1234, 0x0001),
360 length: 16,
361 request_id: 0xABCD_0042,
362 protocol_version: 0x01,
363 interface_version: 0x03,
364 message_type: MessageTypeField::try_from(0x00).unwrap(), return_code: ReturnCode::Ok,
366 }
367 }
368
369 fn encode_header(h: &Header) -> [u8; 16] {
370 let mut buf = [0u8; 16];
371 h.encode(&mut buf.as_mut_slice()).unwrap();
372 buf
373 }
374
375 #[test]
378 fn upper_header_bytes_layout() {
379 let h = make_header();
380 let ub = h.upper_header_bytes();
381 let rid = h.request_id().to_be_bytes();
382 assert_eq!(ub[0..4], rid);
383 assert_eq!(ub[4], h.protocol_version());
384 assert_eq!(ub[5], h.interface_version());
385 assert_eq!(ub[6], u8::from(h.message_type()));
386 assert_eq!(ub[7], u8::from(h.return_code()));
387 }
388
389 #[test]
392 fn new_sd_fields() {
393 let h = Header::new_sd(0x0000_0001, 28);
394 assert_eq!(h.message_id(), MessageId::SD);
395 assert_eq!(h.length(), 8 + 28);
396 assert_eq!(h.request_id(), 0x0000_0001);
397 assert_eq!(h.protocol_version(), 0x01);
398 assert_eq!(h.interface_version(), 0x01);
399 assert_eq!(h.return_code(), ReturnCode::Ok);
400 }
401
402 #[test]
405 fn is_sd_true_for_sd_header() {
406 let h = Header::new_sd(0, 12);
407 assert!(h.is_sd());
408 }
409
410 #[test]
411 fn is_sd_false_for_non_sd_header() {
412 let h = make_header();
413 assert!(!h.is_sd());
414 }
415
416 #[test]
419 fn payload_size_returns_length_minus_8() {
420 let h = Header {
421 length: 24,
422 ..make_header()
423 };
424 assert_eq!(h.payload_size(), 16);
425 }
426
427 #[test]
430 fn set_request_id_updates_value() {
431 let mut h = make_header();
432 h.set_request_id(0xDEAD_BEEF);
433 assert_eq!(h.request_id(), 0xDEAD_BEEF);
434 }
435
436 #[test]
439 fn required_size_is_16() {
440 assert_eq!(make_header().required_size(), 16);
441 }
442
443 #[test]
446 fn encode_parse_round_trip() {
447 let h = make_header();
448 let buf = encode_header(&h);
449 let (view, remaining) = HeaderView::parse(&buf[..]).unwrap();
450 assert_eq!(view.to_owned(), h);
451 assert!(remaining.is_empty());
452 }
453
454 #[test]
455 fn encode_returns_16() {
456 let h = make_header();
457 let mut buf = [0u8; 16];
458 let n = h.encode(&mut buf.as_mut_slice()).unwrap();
459 assert_eq!(n, 16);
460 }
461
462 #[test]
463 fn sd_header_round_trips() {
464 let h = Header::new_sd(0x0000_0042, 28);
465 let buf = encode_header(&h);
466 let (view, _) = HeaderView::parse(&buf[..]).unwrap();
467 assert_eq!(view.to_owned(), h);
468 }
469
470 #[test]
473 fn parse_exact_size_slice_returns_empty_remainder() {
474 let h = make_header();
475 let buf = encode_header(&h);
476 let (view, remaining) = HeaderView::parse(&buf).unwrap();
478 assert_eq!(view.to_owned(), h);
479 assert!(remaining.is_empty());
480 }
481
482 #[test]
485 fn parse_invalid_protocol_version_returns_error() {
486 let mut h = make_header();
487 h.protocol_version = 0x02;
488 let mid = h.message_id.message_id().to_be_bytes();
490 let len = h.length.to_be_bytes();
491 let rid = h.request_id.to_be_bytes();
492 let buf: [u8; 16] = [
493 mid[0], mid[1], mid[2], mid[3], len[0], len[1], len[2], len[3], rid[0], rid[1], rid[2],
494 rid[3], 0x02, 0x03, 0x00, 0x00,
496 ];
497 assert!(matches!(
498 HeaderView::parse(&buf[..]),
499 Err(Error::InvalidProtocolVersion(0x02))
500 ));
501 }
502
503 #[test]
504 fn parse_invalid_message_type_returns_error() {
505 let h = make_header();
506 let mut buf = encode_header(&h);
507 buf[14] = 0xFF; assert!(matches!(
509 HeaderView::parse(&buf[..]),
510 Err(Error::InvalidMessageTypeField(0xFF))
511 ));
512 }
513
514 #[test]
515 fn parse_invalid_return_code_returns_error() {
516 let h = make_header();
517 let mut buf = encode_header(&h);
518 buf[15] = 0x5F; assert!(matches!(
520 HeaderView::parse(&buf[..]),
521 Err(Error::InvalidReturnCode(0x5F))
522 ));
523 }
524
525 #[test]
526 fn parse_truncated_input_returns_eof() {
527 let buf: [u8; 4] = [0x00, 0x00, 0x00, 0x00];
528 assert!(matches!(
529 HeaderView::parse(&buf[..]),
530 Err(Error::UnexpectedEof)
531 ));
532 }
533
534 #[test]
537 fn from_fields_round_trip() {
538 let h = make_header();
539 let h2 = Header::from_fields(
540 h.message_id(),
541 h.length(),
542 h.request_id(),
543 h.protocol_version(),
544 h.interface_version(),
545 h.message_type(),
546 h.return_code(),
547 );
548 assert_eq!(h, h2);
549 }
550
551 #[test]
554 fn new_event_fields() {
555 let h = Header::new_event(0x5B, 0x8001, 0x0001, 0x01, 0x03, 10);
556 assert_eq!(h.message_id().service_id(), 0x5B);
557 assert_eq!(h.message_id().method_id(), 0x8001);
558 assert_eq!(h.request_id(), 0x0001);
559 assert_eq!(h.protocol_version(), 0x01);
560 assert_eq!(h.interface_version(), 0x03);
561 assert_eq!(h.length(), 18); assert_eq!(h.return_code(), ReturnCode::Ok);
563 }
564
565 #[test]
568 fn new_constructor_sets_length() {
569 let h = Header::new(
570 MessageId::new_from_service_and_method(0x1234, 0x0001),
571 0x0001,
572 0x01,
573 0x01,
574 MessageTypeField::try_from(0x00).unwrap(),
575 ReturnCode::Ok,
576 100,
577 );
578 assert_eq!(h.length(), 108); assert_eq!(h.payload_size(), 100);
580 }
581
582 #[test]
585 fn header_view_accessors() {
586 let h = make_header();
587 let buf = encode_header(&h);
588 let (view, _) = HeaderView::parse(&buf[..]).unwrap();
589 assert_eq!(view.message_id(), h.message_id());
590 assert_eq!(view.length(), h.length());
591 assert_eq!(view.request_id(), h.request_id());
592 assert_eq!(view.payload_size(), h.payload_size());
593 assert_eq!(view.protocol_version(), h.protocol_version());
594 assert_eq!(view.interface_version(), h.interface_version());
595 assert_eq!(view.message_type(), h.message_type());
596 assert_eq!(view.return_code(), h.return_code());
597 assert_eq!(view.is_sd(), h.is_sd());
598 }
599
600 #[test]
603 fn encode_to_slice_works() {
604 let h = make_header();
605 let mut buf = [0u8; 16];
606 let n = h.encode_to_slice(&mut buf).unwrap();
607 assert_eq!(n, 16);
608 let (view, _) = HeaderView::parse(&buf).unwrap();
609 assert_eq!(view.to_owned(), h);
610 }
611
612 #[cfg(feature = "std")]
613 #[test]
614 fn encode_to_vec_works() {
615 let h = make_header();
616 let buf = h.encode_to_vec().unwrap();
617 assert_eq!(buf.len(), 16);
618 let (view, _) = HeaderView::parse(&buf).unwrap();
619 assert_eq!(view.to_owned(), h);
620 }
621}