1use std::num::NonZeroI32;
8
9use log::warn;
10use net_types::ip::{Ip, Ipv4Addr, Ipv6Addr};
11use netlink_packet_core::buffer::NETLINK_HEADER_LEN;
12use netlink_packet_core::constants::{
13 NLM_F_ACK, NLM_F_APPEND, NLM_F_ATOMIC, NLM_F_CREATE, NLM_F_DUMP, NLM_F_ECHO, NLM_F_EXCL,
14 NLM_F_MATCH, NLM_F_MULTIPART, NLM_F_REPLACE, NLM_F_REQUEST, NLM_F_ROOT,
15};
16use netlink_packet_core::{
17 DoneMessage, ErrorMessage, NetlinkHeader, NetlinkMessage, NetlinkPayload, NetlinkSerializable,
18};
19use netlink_packet_route::route::RouteAddress;
20use netlink_packet_utils::Emitable as _;
21
22use crate::netlink_packet::errno::Errno;
23
24pub(crate) const UNSPECIFIED_SEQUENCE_NUMBER: u32 = 0;
25
26const DONE_ERROR_CODE: i32 = 0;
28
29pub(crate) fn new_done<T: NetlinkSerializable>(req_header: NetlinkHeader) -> NetlinkMessage<T> {
31 let mut done = DoneMessage::default();
32 done.code = DONE_ERROR_CODE;
33 let payload = NetlinkPayload::<T>::Done(done);
34 let mut resp_header = NetlinkHeader::default();
35 resp_header.sequence_number = req_header.sequence_number;
36 resp_header.flags |= NLM_F_MULTIPART;
37 let mut message = NetlinkMessage::new(resp_header, payload);
38 message.finalize();
40 message
41}
42
43pub(crate) fn ip_addr_from_route<I: Ip>(route_addr: &RouteAddress) -> Result<I::Addr, Errno> {
45 I::map_ip(
46 (),
47 |()| match route_addr {
48 RouteAddress::Inet(v4_addr) => Ok(Ipv4Addr::new(v4_addr.octets())),
49 RouteAddress::Inet6(_) => {
50 warn!("expected IPv4 address from route but got an IPv6 address");
51 Err(Errno::EINVAL)
52 }
53 RouteAddress::Mpls(_) | RouteAddress::Other(_) | _ => Err(Errno::ENOTSUP),
54 },
55 |()| match route_addr {
56 RouteAddress::Inet6(v6_addr) => Ok(Ipv6Addr::new(v6_addr.segments())),
57 RouteAddress::Inet(_) => {
58 warn!("expected IPv6 address from route but got an IPv4 address");
59 Err(Errno::EINVAL)
60 }
61 RouteAddress::Mpls(_) | RouteAddress::Other(_) | _ => Err(Errno::ENOTSUP),
62 },
63 )
64}
65
66pub(crate) mod errno {
67 use net_types::ip::GenericOverIp;
68
69 use super::*;
70
71 #[derive(Copy, Clone, Debug, PartialEq, GenericOverIp)]
75 #[generic_over_ip()]
76 pub struct Errno(i32);
77
78 impl Errno {
79 pub(crate) const EADDRNOTAVAIL: Errno = Errno::new(libc::EADDRNOTAVAIL).unwrap();
80 pub(crate) const EAFNOSUPPORT: Errno = Errno::new(libc::EAFNOSUPPORT).unwrap();
81 pub(crate) const EBUSY: Errno = Errno::new(libc::EBUSY).unwrap();
82 pub(crate) const EEXIST: Errno = Errno::new(libc::EEXIST).unwrap();
83 pub(crate) const EINVAL: Errno = Errno::new(libc::EINVAL).unwrap();
84 pub(crate) const ENODEV: Errno = Errno::new(libc::ENODEV).unwrap();
85 pub(crate) const ENOENT: Errno = Errno::new(libc::ENOENT).unwrap();
86 pub(crate) const ENOTSUP: Errno = Errno::new(libc::ENOTSUP).unwrap();
87 pub(crate) const ESRCH: Errno = Errno::new(libc::ESRCH).unwrap();
88 pub(crate) const ETOOMANYREFS: Errno = Errno::new(libc::ETOOMANYREFS).unwrap();
89 pub(crate) const ENOBUFS: Errno = Errno::new(libc::ENOBUFS).unwrap();
90
91 pub const fn new(code: i32) -> Option<Self> {
95 if code.is_positive() { Some(Errno(code)) } else { None }
96 }
97 }
98
99 impl From<Errno> for NonZeroI32 {
100 fn from(Errno(code): Errno) -> Self {
101 NonZeroI32::new(code).expect("Errno's code must be non-zero")
102 }
103 }
104
105 impl From<Errno> for i32 {
106 fn from(Errno(code): Errno) -> Self {
107 code
108 }
109 }
110
111 #[cfg(test)]
112 mod tests {
113 use super::*;
114 use test_case::test_case;
115
116 #[test_case(i32::MIN, None; "min")]
117 #[test_case(-10, None; "negative")]
118 #[test_case(0, None; "zero")]
119 #[test_case(10, Some(10); "positive")]
120 #[test_case(i32::MAX, Some(i32::MAX); "max")]
121 fn test_new_errno(raw_code: i32, expected_code: Option<i32>) {
122 assert_eq!(Errno::new(raw_code).map(Into::<i32>::into), expected_code)
123 }
124 }
125}
126
127pub(crate) fn new_error<T: NetlinkSerializable>(
131 error: Result<(), errno::Errno>,
132 req_header: NetlinkHeader,
133) -> NetlinkMessage<T> {
134 let error = {
135 assert_eq!(req_header.buffer_len(), NETLINK_HEADER_LEN);
136 let mut buffer = vec![0; NETLINK_HEADER_LEN];
137 req_header.emit(&mut buffer);
138
139 let code = match error {
140 Ok(()) => None,
141
142 Err(e) => Some(-NonZeroI32::from(e)),
144 };
145
146 let mut error = ErrorMessage::default();
147 error.code = code;
148 error.header = buffer;
149 error
150 };
151
152 let payload = NetlinkPayload::<T>::Error(error);
153 let mut resp_header = NetlinkHeader::default();
156 resp_header.sequence_number = req_header.sequence_number;
157 let mut message = NetlinkMessage::new(resp_header, payload);
158 message.finalize();
160 message
161}
162
163#[derive(Clone, Copy, Debug, PartialEq)]
165pub(crate) enum NetlinkRequestType {
166 New,
168 Get,
170 Set,
172 Del,
174}
175
176pub(crate) fn netlink_flags_debug_string(flags: u16, request_type: NetlinkRequestType) -> String {
181 let mut flags_dbg = vec![];
182 if (flags & NLM_F_REQUEST) == NLM_F_REQUEST {
183 flags_dbg.push("REQUEST");
184 }
185 if (flags & NLM_F_MULTIPART) == NLM_F_MULTIPART {
186 flags_dbg.push("MULTI");
187 }
188 if (flags & NLM_F_ACK) == NLM_F_ACK {
189 flags_dbg.push("ACK");
190 }
191 if (flags & NLM_F_ECHO) == NLM_F_ECHO {
192 flags_dbg.push("ECHO");
193 }
194 match request_type {
195 NetlinkRequestType::Get => {
196 if (flags & NLM_F_DUMP) == NLM_F_DUMP {
197 flags_dbg.push("DUMP");
198 } else {
199 if (flags & NLM_F_ROOT) == NLM_F_ROOT {
201 flags_dbg.push("ROOT");
202 }
203 if (flags & NLM_F_MATCH) == NLM_F_MATCH {
204 flags_dbg.push("MATCH");
205 }
206 }
207 if (flags & NLM_F_ATOMIC) == NLM_F_ATOMIC {
208 flags_dbg.push("ATOMIC");
209 }
210 }
211 NetlinkRequestType::New => {
212 if (flags & NLM_F_REPLACE) == NLM_F_REPLACE {
213 flags_dbg.push("REPLACE");
214 }
215 if (flags & NLM_F_EXCL) == NLM_F_EXCL {
216 flags_dbg.push("EXCL");
217 }
218 if (flags & NLM_F_CREATE) == NLM_F_CREATE {
219 flags_dbg.push("CREATE");
220 }
221 if (flags & NLM_F_APPEND) == NLM_F_APPEND {
222 flags_dbg.push("APPEND");
223 }
224 }
225 NetlinkRequestType::Set | NetlinkRequestType::Del => {}
226 }
227 flags_dbg.join("|")
228}
229
230#[cfg(test)]
231mod tests {
232 use super::*;
233
234 use assert_matches::assert_matches;
235 use netlink_packet_core::{NLMSG_DONE, NLMSG_ERROR, NetlinkBuffer};
236 use netlink_packet_route::RouteNetlinkMessage;
237 use netlink_packet_utils::Parseable as _;
238 use test_case::test_case;
239
240 use crate::netlink_packet::errno::Errno;
241
242 #[test_case(0, Ok(()); "ACK")]
243 #[test_case(0, Err(Errno::EINVAL); "EINVAL")]
244 #[test_case(1, Err(Errno::ENODEV); "ENODEV")]
245 fn test_new_error(sequence_number: u32, expected_error: Result<(), Errno>) {
246 let mut expected_header = NetlinkHeader::default();
248 expected_header.length = 0x01234567;
249 expected_header.message_type = 0x89AB;
250 expected_header.flags = 0xCDEF;
251 expected_header.sequence_number = sequence_number;
252 expected_header.port_number = 0x00000000;
253
254 let error = new_error::<RouteNetlinkMessage>(expected_error, expected_header);
255 let mut buf = vec![0; error.buffer_len()];
257 error.serialize(&mut buf);
258
259 let (header, payload) = error.into_parts();
260 assert_eq!(header.message_type, NLMSG_ERROR);
261 assert_eq!(header.sequence_number, sequence_number);
262 assert_matches!(
263 payload,
264 NetlinkPayload::Error(ErrorMessage{ code, header, .. }) => {
265 let expected_code = match expected_error {
266 Ok(()) => None,
267 Err(e) => Some(-NonZeroI32::from(e)),
268 };
269 assert_eq!(code, expected_code);
270 assert_eq!(
271 NetlinkHeader::parse(&NetlinkBuffer::new_unchecked(&header)).unwrap(),
274 expected_header,
275 );
276 }
277 );
278 }
279
280 #[test_case(0; "seq_0")]
281 #[test_case(1; "seq_1")]
282 fn test_new_done(sequence_number: u32) {
283 let mut req_header = NetlinkHeader::default();
284 req_header.sequence_number = sequence_number;
285
286 let done = new_done::<RouteNetlinkMessage>(req_header);
287 let mut buf = vec![0; done.buffer_len()];
289 done.serialize(&mut buf);
290
291 let (header, payload) = done.into_parts();
292 assert_eq!(header.sequence_number, sequence_number);
293 assert_eq!(header.message_type, NLMSG_DONE);
294 assert_eq!(header.flags, NLM_F_MULTIPART);
295 assert_matches!(
296 payload,
297 NetlinkPayload::Done(DoneMessage {code, extended_ack, ..}) => {
298 assert_eq!(code, DONE_ERROR_CODE);
299 assert_eq!(extended_ack, Vec::<u8>::new());
300 }
301 );
302 }
303
304 #[test_case(
305 0,
306 NetlinkRequestType::Get => "";
307 "no flags"
308 )]
309 #[test_case(
310 NLM_F_REQUEST,
311 NetlinkRequestType::Get => "REQUEST";
312 "request only"
313 )]
314 #[test_case(
315 NLM_F_REQUEST|NLM_F_MULTIPART|NLM_F_ACK|NLM_F_ECHO,
316 NetlinkRequestType::Get => "REQUEST|MULTI|ACK|ECHO";
317 "all generic flags"
318 )]
319 #[test_case(
320 NLM_F_REQUEST|NLM_F_DUMP,
321 NetlinkRequestType::Get => "REQUEST|DUMP";
322 "dump request"
323 )]
324 #[test_case(
325 NLM_F_REQUEST|NLM_F_MATCH|NLM_F_ROOT,
326 NetlinkRequestType::Get => "REQUEST|DUMP";
327 "dump is alias for match|root"
328 )]
329 #[test_case(
330 NLM_F_REQUEST|NLM_F_ATOMIC,
331 NetlinkRequestType::Get => "REQUEST|ATOMIC";
332 "other Get flags"
333 )]
334 #[test_case(
335 NLM_F_REQUEST|NLM_F_REPLACE|NLM_F_EXCL|NLM_F_CREATE|NLM_F_APPEND,
336 NetlinkRequestType::New
337 => "REQUEST|REPLACE|EXCL|CREATE|APPEND";
338 "New flags"
339 )]
340 #[test_case(
341 NLM_F_REQUEST|NLM_F_REPLACE|NLM_F_EXCL|NLM_F_CREATE|NLM_F_APPEND,
342 NetlinkRequestType::Del => "REQUEST";
343 "type-inappropriate flags ignored"
344 )]
345 fn netlink_flags_debug_string_tests(flags: u16, request_type: NetlinkRequestType) -> String {
346 netlink_flags_debug_string(flags, request_type)
347 }
348}