Skip to main content

netlink_packet_core/
message.rs

1// SPDX-License-Identifier: MIT
2
3use std::fmt::Debug;
4
5use crate::{
6    done::DONE_HEADER_LEN,
7    payload::{NLMSG_DONE, NLMSG_ERROR, NLMSG_NOOP, NLMSG_OVERRUN},
8    DecodeError, DoneBuffer, DoneMessage, Emitable, ErrorBuffer, ErrorContext,
9    ErrorMessage, NetlinkBuffer, NetlinkDeserializable, NetlinkHeader,
10    NetlinkPayload, NetlinkSerializable, Parseable,
11};
12
13/// Represent a netlink message.
14#[derive(Debug, PartialEq, Eq, Clone)]
15#[non_exhaustive]
16pub struct NetlinkMessage<I> {
17    /// Message header (this is common to all the netlink protocols)
18    pub header: NetlinkHeader,
19    /// Inner message, which depends on the netlink protocol being used.
20    pub payload: NetlinkPayload<I>,
21}
22
23impl<I> NetlinkMessage<I> {
24    /// Create a new netlink message from the given header and payload
25    pub fn new(header: NetlinkHeader, payload: NetlinkPayload<I>) -> Self {
26        NetlinkMessage { header, payload }
27    }
28
29    /// Consume this message and return its header and payload
30    pub fn into_parts(self) -> (NetlinkHeader, NetlinkPayload<I>) {
31        (self.header, self.payload)
32    }
33}
34
35impl<I> NetlinkMessage<I>
36where
37    I: NetlinkDeserializable,
38{
39    /// Parse the given buffer as a netlink message
40    pub fn deserialize(buffer: &[u8]) -> Result<Self, DecodeError> {
41        let netlink_buffer = NetlinkBuffer::new_checked(&buffer)
42            .context("failed deserializing NetlinkMessage")?;
43        <Self as Parseable<NetlinkBuffer<&&[u8]>>>::parse(&netlink_buffer)
44    }
45}
46
47impl<I> NetlinkMessage<I>
48where
49    I: NetlinkSerializable,
50{
51    /// Return the length of this message in bytes
52    pub fn buffer_len(&self) -> usize {
53        <Self as Emitable>::buffer_len(self)
54    }
55
56    /// Serialize this message and write the serialized data into the
57    /// given buffer. `buffer` must big large enough for the whole
58    /// message to fit, otherwise, this method will panic. To know how
59    /// big the serialized message is, call `buffer_len()`.
60    ///
61    /// # Panic
62    ///
63    /// This method panics if the buffer is not big enough.
64    pub fn serialize(&self, buffer: &mut [u8]) {
65        self.emit(buffer)
66    }
67
68    /// Ensure the header (`NetlinkHeader`) is consistent with the payload
69    /// (`NetlinkPayload`):
70    ///
71    /// - compute the payload length and set the header's length field
72    /// - check the payload type and set the header's message type field
73    ///   accordingly
74    ///
75    /// If you are not 100% sure the header is correct, this method should be
76    /// called before calling [`Emitable::emit()`](trait.Emitable.html#
77    /// tymethod.emit), as it could panic if the header is inconsistent with
78    /// the rest of the message.
79    pub fn finalize(&mut self) {
80        self.header.length = self.buffer_len() as u32;
81        self.header.message_type = self.payload.message_type();
82    }
83}
84
85impl<B, I> Parseable<NetlinkBuffer<&B>> for NetlinkMessage<I>
86where
87    B: AsRef<[u8]>,
88    I: NetlinkDeserializable,
89{
90    fn parse(buf: &NetlinkBuffer<&B>) -> Result<Self, DecodeError> {
91        use self::NetlinkPayload::*;
92
93        let header =
94            <NetlinkHeader as Parseable<NetlinkBuffer<&B>>>::parse(buf)
95                .context("failed parsing NetlinkHeader")?;
96
97        let bytes = buf.payload();
98        let payload = match header.message_type {
99            NLMSG_ERROR => {
100                let msg = ErrorBuffer::new_checked(&bytes)
101                    .and_then(|buf| ErrorMessage::parse(&buf))
102                    .context("failed parsing NLMSG_ERROR")?;
103                Error(msg)
104            }
105            NLMSG_NOOP => Noop,
106            NLMSG_DONE => {
107                // Linux kernel allows zero sized of NLMSG_DONE
108                let msg = if bytes.is_empty() {
109                    DoneBuffer::new_checked(&[0u8; DONE_HEADER_LEN])
110                        .and_then(|buf| DoneMessage::parse(&buf))
111                        .context("failed to parse NLMSG_DONE")?
112                } else {
113                    DoneBuffer::new_checked(&bytes)
114                        .and_then(|buf| DoneMessage::parse(&buf))
115                        .context("failed to parse NLMSG_DONE")?
116                };
117                Done(msg)
118            }
119            NLMSG_OVERRUN => Overrun(bytes.to_vec()),
120            message_type => match I::deserialize(&header, bytes) {
121                Err(e) => {
122                    return Err(format!(
123                        "Failed to parse message with type {message_type}: {e}"
124                    )
125                    .into())
126                }
127                Ok(inner_msg) => InnerMessage(inner_msg),
128            },
129        };
130        Ok(NetlinkMessage { header, payload })
131    }
132}
133
134impl<I> Emitable for NetlinkMessage<I>
135where
136    I: NetlinkSerializable,
137{
138    fn buffer_len(&self) -> usize {
139        use self::NetlinkPayload::*;
140
141        let payload_len = match self.payload {
142            Noop => 0,
143            Done(ref msg) => msg.buffer_len(),
144            Overrun(ref bytes) => bytes.len(),
145            Error(ref msg) => msg.buffer_len(),
146            InnerMessage(ref msg) => msg.buffer_len(),
147        };
148
149        self.header.buffer_len() + payload_len
150    }
151
152    fn emit(&self, buffer: &mut [u8]) {
153        use self::NetlinkPayload::*;
154
155        self.header.emit(buffer);
156
157        let buffer =
158            &mut buffer[self.header.buffer_len()..self.header.length as usize];
159        match self.payload {
160            Noop => {}
161            Done(ref msg) => msg.emit(buffer),
162            Overrun(ref bytes) => buffer.copy_from_slice(bytes),
163            Error(ref msg) => msg.emit(buffer),
164            InnerMessage(ref msg) => msg.serialize(buffer),
165        }
166    }
167}
168
169impl<T> From<T> for NetlinkMessage<T>
170where
171    T: Into<NetlinkPayload<T>>,
172{
173    fn from(inner_message: T) -> Self {
174        NetlinkMessage {
175            header: NetlinkHeader::default(),
176            payload: inner_message.into(),
177        }
178    }
179}
180
181// test data are using hard coded little endian byte order, not for big-endian
182#[cfg(not(target_endian = "big"))]
183#[cfg(test)]
184mod tests {
185    use super::*;
186
187    use std::{convert::Infallible, mem::size_of, num::NonZeroI32};
188
189    #[derive(Clone, Debug, Default, PartialEq)]
190    struct FakeNetlinkInnerMessage;
191
192    impl NetlinkSerializable for FakeNetlinkInnerMessage {
193        fn message_type(&self) -> u16 {
194            unimplemented!("unused by tests")
195        }
196
197        fn buffer_len(&self) -> usize {
198            unimplemented!("unused by tests")
199        }
200
201        fn serialize(&self, _buffer: &mut [u8]) {
202            unimplemented!("unused by tests")
203        }
204    }
205
206    impl NetlinkDeserializable for FakeNetlinkInnerMessage {
207        type Error = Infallible;
208
209        fn deserialize(
210            _header: &NetlinkHeader,
211            _payload: &[u8],
212        ) -> Result<Self, Self::Error> {
213            unimplemented!("unused by tests")
214        }
215    }
216
217    #[test]
218    fn test_done() {
219        let header = NetlinkHeader::default();
220        let done_msg = DoneMessage {
221            code: 0,
222            extended_ack: vec![6, 7, 8, 9],
223        };
224        let mut want = NetlinkMessage::new(
225            header,
226            NetlinkPayload::<FakeNetlinkInnerMessage>::Done(done_msg.clone()),
227        );
228        want.finalize();
229
230        let len = want.buffer_len();
231        assert_eq!(
232            len,
233            header.buffer_len()
234                + size_of::<i32>()
235                + done_msg.extended_ack.len()
236        );
237
238        let mut buf = vec![1; len];
239        want.emit(&mut buf);
240
241        let done_buf = DoneBuffer::new(&buf[header.buffer_len()..]);
242        assert_eq!(done_buf.code(), done_msg.code);
243        assert_eq!(done_buf.extended_ack(), &done_msg.extended_ack);
244
245        let got = NetlinkMessage::parse(&NetlinkBuffer::new(&buf)).unwrap();
246        assert_eq!(got, want);
247    }
248
249    #[test]
250    fn test_error() {
251        // SAFETY: value is non-zero.
252        const ERROR_CODE: NonZeroI32 = NonZeroI32::new(-8765).unwrap();
253
254        let header = NetlinkHeader::default();
255        let error_msg = ErrorMessage {
256            code: Some(ERROR_CODE),
257            header: vec![],
258        };
259        let mut want = NetlinkMessage::new(
260            header,
261            NetlinkPayload::<FakeNetlinkInnerMessage>::Error(error_msg.clone()),
262        );
263        want.finalize();
264
265        let len = want.buffer_len();
266        assert_eq!(len, header.buffer_len() + error_msg.buffer_len());
267
268        let mut buf = vec![1; len];
269        want.emit(&mut buf);
270
271        let error_buf = ErrorBuffer::new(&buf[header.buffer_len()..]);
272        assert_eq!(error_buf.code(), error_msg.code);
273
274        let got = NetlinkMessage::parse(&NetlinkBuffer::new(&buf)).unwrap();
275        assert_eq!(got, want);
276    }
277}