1use netlink_packet_core::NetlinkMessage;
4use netlink_packet_generic::GenlMessage;
5use netlink_packet_wireguard::WireguardMessage;
6
7#[derive(Clone, Copy, Eq, PartialEq, Debug)]
12#[non_exhaustive]
13pub enum ErrorKind {
14 Bug,
15 NetlinkError,
16 DecodeError,
17 InvalidKey,
19 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 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 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 let message = request_message();
164 let mut raw = vec![0u8; 16 + 4 + message.buffer_len()];
165 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 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}