Skip to main content

nl_wireguard/
error.rs

1// SPDX-License-Identifier: MIT
2
3use netlink_packet_core::NetlinkMessage;
4use netlink_packet_generic::GenlMessage;
5use netlink_packet_wireguard::WireguardMessage;
6
7/// Kind of a [WireguardError].
8///
9/// This enum is `#[non_exhaustive]`, a `match` on it needs a wildcard arm.
10/// Otherwise every new error kind would break the code of every user.
11#[derive(Clone, Copy, Eq, PartialEq, Debug)]
12#[non_exhaustive]
13pub enum ErrorKind {
14    Bug,
15    NetlinkError,
16    DecodeError,
17    /// Invalid key, should be base64 encoded of [u8; 32]
18    InvalidKey,
19    /// Invalid input from the caller, e.g. missing required property
20    InvalidInput,
21}
22
23impl std::fmt::Display for ErrorKind {
24    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
25        write!(
26            f,
27            "{}",
28            match self {
29                Self::Bug => "bug",
30                Self::NetlinkError => "netlink_error",
31                Self::DecodeError => "decode_error",
32                Self::InvalidKey => "invalid_key",
33                Self::InvalidInput => "invalid_input",
34            }
35        )
36    }
37}
38
39#[derive(Clone, Eq, PartialEq, Debug)]
40pub struct WireguardError {
41    pub kind: ErrorKind,
42    pub msg: String,
43    pub netlink_msg: Option<NetlinkMessage<GenlMessage<WireguardMessage>>>,
44}
45
46impl std::fmt::Display for WireguardError {
47    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
48        if let Some(nl_msg) = self.netlink_msg.as_ref() {
49            write!(
50                f,
51                "{}: {}, netlink message: {:?}",
52                self.kind, self.msg, nl_msg
53            )
54        } else {
55            write!(f, "{}: {}", self.kind, self.msg)
56        }
57    }
58}
59
60impl std::error::Error for WireguardError {}
61
62impl WireguardError {
63    /// Create a new [WireguardError].
64    ///
65    /// Key material is replaced with zeros in `netlink_msg`, so that keys
66    /// which the kernel echoed back can not end up in logs.
67    pub fn new(
68        kind: ErrorKind,
69        msg: String,
70        netlink_msg: Option<NetlinkMessage<GenlMessage<WireguardMessage>>>,
71    ) -> Self {
72        Self {
73            kind,
74            msg,
75            netlink_msg: netlink_msg.map(|mut netlink_msg| {
76                crate::redact::redact_message(&mut netlink_msg);
77                netlink_msg
78            }),
79        }
80    }
81}
82
83#[cfg(test)]
84mod tests {
85    use std::num::NonZeroI32;
86
87    use netlink_packet_core::{
88        Emitable, ErrorMessage, NetlinkHeader, NetlinkPayload,
89    };
90    use netlink_packet_wireguard::{
91        WireguardAttribute, WireguardCmd, WireguardPeer, WireguardPeerAttribute,
92    };
93
94    use super::*;
95
96    const PRIVATE_KEY: [u8; 32] = [0x7b; 32];
97    const PRESHARED_KEY: [u8; 32] = [0x5a; 32];
98    const PUBLIC_KEY: [u8; 32] = [0x11; 32];
99
100    fn request_message() -> WireguardMessage {
101        WireguardMessage {
102            cmd: WireguardCmd::SetDevice,
103            attributes: vec![
104                WireguardAttribute::IfName("wg0".to_string()),
105                WireguardAttribute::PrivateKey(PRIVATE_KEY),
106                WireguardAttribute::Peers(vec![WireguardPeer(vec![
107                    WireguardPeerAttribute::PublicKey(PUBLIC_KEY),
108                    WireguardPeerAttribute::PresharedKey(PRESHARED_KEY),
109                ])]),
110            ],
111        }
112    }
113
114    /// Assert that all keys of `attributes` are zeroed and report whether
115    /// a private and a preshared key was found.
116    fn assert_redacted(attributes: &[WireguardAttribute]) -> (bool, bool) {
117        let mut has_private_key = false;
118        let mut has_preshared_key = false;
119        for attribute in attributes {
120            match attribute {
121                WireguardAttribute::PrivateKey(key) => {
122                    has_private_key = true;
123                    assert_eq!(*key, [0u8; 32]);
124                }
125                WireguardAttribute::Peers(peers) => {
126                    for peer in peers {
127                        for attribute in &peer.0 {
128                            if let WireguardPeerAttribute::PresharedKey(key) =
129                                attribute
130                            {
131                                has_preshared_key = true;
132                                assert_eq!(*key, [0u8; 32]);
133                            }
134                        }
135                    }
136                }
137                _ => (),
138            }
139        }
140        (has_private_key, has_preshared_key)
141    }
142
143    #[test]
144    fn keys_are_redacted_from_request_messages() {
145        let err = WireguardError::new(
146            ErrorKind::NetlinkError,
147            "test".to_string(),
148            Some(NetlinkMessage::from(GenlMessage::from_payload(
149                request_message(),
150            ))),
151        );
152
153        let stored = err.netlink_msg.as_ref().expect("no netlink message");
154        let NetlinkPayload::InnerMessage(genl_msg) = &stored.payload else {
155            panic!("unexpected payload {:?}", stored.payload);
156        };
157        assert_eq!(assert_redacted(&genl_msg.payload.attributes), (true, true));
158    }
159
160    #[test]
161    fn echoed_requests_are_dropped() {
162        // The kernel echoes the raw request in its error reply.
163        let message = request_message();
164        let mut raw = vec![0u8; 16 + 4 + message.buffer_len()];
165        // `cmd` and `version` of the generic netlink header.
166        raw[16] = 1;
167        raw[17] = 1;
168        message.emit(&mut raw[20..]);
169
170        let mut error_message = ErrorMessage::default();
171        error_message.code = NonZeroI32::new(-22);
172        error_message.header = raw;
173
174        let err = WireguardError::new(
175            ErrorKind::NetlinkError,
176            "test".to_string(),
177            Some(NetlinkMessage::new(
178                NetlinkHeader::default(),
179                NetlinkPayload::Error(error_message),
180            )),
181        );
182
183        let stored = err.netlink_msg.as_ref().expect("no netlink message");
184        let NetlinkPayload::Error(error_message) = &stored.payload else {
185            panic!("unexpected payload {:?}", stored.payload);
186        };
187        // The raw copy of the request is dropped instead of being redacted
188        // in place, a key can not be missed that way.
189        assert!(error_message.header.is_empty());
190        assert!(!format!("{err:?}").contains(&hex(&PRIVATE_KEY)));
191        assert!(!format!("{err:?}").contains(&hex(&PRESHARED_KEY)));
192    }
193
194    fn hex(key: &[u8]) -> String {
195        key.iter()
196            .map(|byte| byte.to_string())
197            .collect::<Vec<String>>()
198            .join(", ")
199    }
200}