1use std::io;
2#[cfg(unix)]
3use std::os::unix::fs::PermissionsExt;
4use std::path::Path;
5use std::process::Command;
6use std::sync::LazyLock;
7use std::sync::atomic::{AtomicU32, Ordering};
8
9use parking_lot::RwLock;
10
11use base64::Engine;
12use base64::engine::general_purpose::URL_SAFE_NO_PAD;
13use rand::RngExt;
14use ring::digest::{SHA256, digest};
15
16use crate::DataLenType;
17
18pub type ChecksumType = u32;
19
20pub const ENV_MSG_HEADER_KEY: &str = "MSG_HEADER_KEY";
22pub const MACHINE_MSG_HEADER_KEY_PATH: &str = "/var/lib/pb-mapper-server/msg_header_key";
24pub const TEMP_CREDENTIAL_PREFIX: &str = "pbmt1_";
25pub const ADMIN_KEY_LEN: usize = 32;
26
27pub fn is_env_safe_admin_key(bytes: &[u8]) -> bool {
30 bytes.len() == ADMIN_KEY_LEN && bytes.iter().all(|byte| byte.is_ascii_graphic())
31}
32
33pub fn env_safe_admin_key_error() -> String {
34 format!(
35 "`{ENV_MSG_HEADER_KEY}` administrator key must be 32 printable ASCII bytes without whitespace or NUL"
36 )
37}
38
39const DERIVE_MSG_HEADER_KEY_TAG: &str = "pb-mapper-msg-header-key-v1";
40const DERIVE_MSG_HEADER_KEY_CHARSET: &[u8] =
41 b"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ";
42
43struct MsgHeaderKeyState {
44 credential: RwLock<Option<Credential>>,
45 load_error: RwLock<Option<String>>,
46 hash: AtomicU32,
47}
48
49#[derive(Clone, Copy, PartialEq, Eq)]
50pub enum Credential {
51 Admin(AesKeyType),
52 Temporary { key_id: u64, key: AesKeyType },
53}
54
55impl std::fmt::Debug for Credential {
56 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
57 match self {
58 Self::Admin(_) => f.debug_tuple("Admin").field(&"[redacted]").finish(),
59 Self::Temporary { key_id, .. } => f
60 .debug_struct("Temporary")
61 .field("key_id", key_id)
62 .field("key", &"[redacted]")
63 .finish(),
64 }
65 }
66}
67
68impl Credential {
69 pub fn key_id(&self) -> u64 {
70 match self {
71 Self::Admin(_) => 0,
72 Self::Temporary { key_id, .. } => *key_id,
73 }
74 }
75
76 pub fn key(&self) -> &AesKeyType {
77 match self {
78 Self::Admin(key) | Self::Temporary { key, .. } => key,
79 }
80 }
81
82 pub fn is_admin(&self) -> bool {
83 matches!(self, Self::Admin(_))
84 }
85}
86
87fn key_len_error(input: &str) -> String {
88 format!(
89 "`{ENV_MSG_HEADER_KEY}` administrator key must be exactly 32 bytes; received {} bytes",
90 input.len()
91 )
92}
93
94fn load_credential_from_env() -> Result<Option<Credential>, String> {
95 let Some(raw) = std::env::var_os(ENV_MSG_HEADER_KEY) else {
96 return Ok(None);
97 };
98 let raw = raw
99 .into_string()
100 .map_err(|_| format!("`{ENV_MSG_HEADER_KEY}` must contain valid UTF-8 credential text"))?;
101 parse_credential(raw.trim()).map(Some)
102}
103
104fn update_runtime_credential(credential: Option<Credential>) {
105 let hash = credential
106 .as_ref()
107 .map(|credential| gen_checksum_by_key(credential.key()))
108 .unwrap_or_default();
109 let mut guard = MSG_HEADER_KEY_STATE.credential.write();
110 *guard = credential;
111 *MSG_HEADER_KEY_STATE.load_error.write() = None;
112 MSG_HEADER_KEY_STATE.hash.store(hash, Ordering::Release);
113}
114
115static MSG_HEADER_KEY_STATE: LazyLock<MsgHeaderKeyState> = LazyLock::new(|| {
120 let (credential, load_error) = match load_credential_from_env() {
121 Ok(credential) => (credential, None),
122 Err(error) => {
123 tracing::error!(reason = "credential_invalid", %error, "invalid MSG_HEADER_KEY");
124 (None, Some(error))
125 }
126 };
127 let hash = credential
128 .as_ref()
129 .map(|credential| gen_checksum_by_key(credential.key()))
130 .unwrap_or_default();
131 MsgHeaderKeyState {
132 credential: RwLock::new(credential),
133 load_error: RwLock::new(load_error),
134 hash: AtomicU32::new(hash),
135 }
136});
137
138pub fn get_process_credential() -> Result<Credential, String> {
140 if let Some(error) = MSG_HEADER_KEY_STATE.load_error.read().clone() {
141 return Err(error);
142 }
143 MSG_HEADER_KEY_STATE.credential.read().ok_or_else(|| {
144 format!("`{ENV_MSG_HEADER_KEY}` is required; no insecure default credential is available")
145 })
146}
147
148pub fn get_msg_header_key() -> Result<Vec<u8>, String> {
150 get_process_credential().map(|credential| credential.key().to_vec())
151}
152
153pub fn set_process_msg_header_key(msg_header_key: Option<&str>) -> Result<(), String> {
158 let normalized = msg_header_key.map(str::trim).unwrap_or("");
159 if normalized.is_empty() {
160 unsafe { std::env::remove_var(ENV_MSG_HEADER_KEY) };
167 update_runtime_credential(None);
168 return Ok(());
169 }
170
171 let credential = parse_credential(normalized)?;
172
173 unsafe { std::env::set_var(ENV_MSG_HEADER_KEY, normalized) };
175 update_runtime_credential(Some(credential));
176 Ok(())
177}
178
179pub fn parse_credential(raw: &str) -> Result<Credential, String> {
180 if let Some(encoded) = raw.strip_prefix(TEMP_CREDENTIAL_PREFIX) {
181 let payload = URL_SAFE_NO_PAD
182 .decode(encoded)
183 .map_err(|_| "temporary credential is not valid base64url".to_string())?;
184 if payload.len() != 45 {
185 return Err(format!(
186 "temporary credential payload must be 45 bytes, got {}",
187 payload.len()
188 ));
189 }
190 if payload[0] != 1 {
191 return Err(format!(
192 "unsupported temporary credential version {}",
193 payload[0]
194 ));
195 }
196 let expected = digest(&SHA256, &payload[..41]);
197 if expected.as_ref()[..4] != payload[41..45] {
198 return Err("temporary credential checksum mismatch".to_string());
199 }
200 let key_id = u64::from_be_bytes(
204 payload[1..9]
205 .try_into()
206 .map_err(|_| "temporary credential key id is malformed".to_string())?,
207 );
208 if key_id == 0 {
209 return Err("temporary credential key id must not be zero".to_string());
210 }
211 let key = payload[9..41]
212 .try_into()
213 .map_err(|_| "temporary credential key is malformed".to_string())?;
214 return Ok(Credential::Temporary { key_id, key });
215 }
216
217 let bytes = raw.as_bytes();
218 if bytes.len() != ADMIN_KEY_LEN {
219 return Err(key_len_error(raw));
220 }
221 if !is_env_safe_admin_key(bytes) {
222 return Err(env_safe_admin_key_error());
223 }
224 Ok(Credential::Admin(
226 bytes.try_into().map_err(|_| key_len_error(raw))?,
227 ))
228}
229
230pub fn encode_temporary_credential(key_id: u64, key: &AesKeyType) -> String {
231 let mut payload = Vec::with_capacity(45);
232 payload.push(1);
233 payload.extend_from_slice(&key_id.to_be_bytes());
234 payload.extend_from_slice(key);
235 let checksum = digest(&SHA256, &payload);
236 payload.extend_from_slice(&checksum.as_ref()[..4]);
237 format!(
238 "{TEMP_CREDENTIAL_PREFIX}{}",
239 URL_SAFE_NO_PAD.encode(payload)
240 )
241}
242
243pub fn setup_machine_msg_header_key() -> io::Result<String> {
256 let hostname = get_machine_hostname()?;
257 let mac_addresses = get_machine_mac_addresses()?;
258 let key = derive_msg_header_key(&hostname, &mac_addresses);
259 set_process_msg_header_key(Some(&key))
260 .map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
261 write_machine_msg_header_key(&key)?;
262 Ok(key)
263}
264
265fn get_machine_hostname() -> io::Result<String> {
266 if let Some(hostname) = normalize_non_empty(std::env::var("HOSTNAME").ok().as_deref()) {
267 return Ok(hostname);
268 }
269
270 if let Ok(content) = std::fs::read_to_string("/etc/hostname")
271 && let Some(hostname) = normalize_non_empty(Some(content.as_str()))
272 {
273 return Ok(hostname);
274 }
275
276 if let Ok(output) = Command::new("hostname").output()
277 && output.status.success()
278 {
279 let stdout = String::from_utf8_lossy(&output.stdout);
280 if let Some(hostname) = normalize_non_empty(Some(stdout.as_ref())) {
281 return Ok(hostname);
282 }
283 }
284
285 Err(io::Error::new(
286 io::ErrorKind::NotFound,
287 "failed to get hostname from HOSTNAME env, /etc/hostname or hostname command",
288 ))
289}
290
291fn normalize_non_empty(input: Option<&str>) -> Option<String> {
292 input.map(str::trim).and_then(|value| {
293 if value.is_empty() {
294 None
295 } else {
296 Some(value.to_ascii_lowercase())
297 }
298 })
299}
300
301fn get_machine_mac_addresses() -> io::Result<Vec<String>> {
302 if let Ok(mac_addresses) = get_machine_mac_addresses_from_sysfs()
303 && !mac_addresses.is_empty()
304 {
305 return Ok(mac_addresses);
306 }
307
308 if let Ok(mac_addresses) = get_machine_mac_addresses_from_ip_link()
309 && !mac_addresses.is_empty()
310 {
311 return Ok(mac_addresses);
312 }
313
314 if let Ok(mac_addresses) = get_machine_mac_addresses_from_ifconfig()
315 && !mac_addresses.is_empty()
316 {
317 return Ok(mac_addresses);
318 }
319
320 Err(io::Error::new(
321 io::ErrorKind::NotFound,
322 "no valid MAC address found from /sys/class/net, `ip link` or `ifconfig`",
323 ))
324}
325
326fn get_machine_mac_addresses_from_sysfs() -> io::Result<Vec<String>> {
327 let mut mac_addresses = Vec::new();
328 for entry in std::fs::read_dir("/sys/class/net")? {
329 let entry = entry?;
330 let interface = match entry.file_name().into_string() {
331 Ok(name) => name,
332 Err(_) => continue,
333 };
334 if interface == "lo" {
335 continue;
336 }
337 let address_path = entry.path().join("address");
338 let mac = match std::fs::read_to_string(address_path) {
339 Ok(mac) => match normalize_mac_address(&mac) {
340 Some(mac) => mac,
341 None => continue,
342 },
343 Err(_) => continue,
344 };
345 mac_addresses.push(format!("{interface}:{mac}"));
346 }
347 normalize_and_validate_mac_entries(&mut mac_addresses);
348 Ok(mac_addresses)
349}
350
351fn get_machine_mac_addresses_from_ip_link() -> io::Result<Vec<String>> {
352 let output = Command::new("ip").arg("link").output()?;
353 if !output.status.success() {
354 return Err(io::Error::other("`ip link` returned non-zero status"));
355 }
356 let mut mac_addresses = Vec::new();
357 let mut current_interface: Option<String> = None;
358 for line in String::from_utf8_lossy(&output.stdout).lines() {
359 if !line.starts_with(' ') {
360 current_interface = parse_interface_name_from_ip_link(line);
361 continue;
362 }
363 let line = line.trim_start();
364 if !line.starts_with("link/ether ") {
365 continue;
366 }
367 let Some(interface) = current_interface.as_ref() else {
368 continue;
369 };
370 let Some(raw_mac) = line.split_whitespace().nth(1) else {
371 continue;
372 };
373 let Some(mac) = normalize_mac_address(raw_mac) else {
374 continue;
375 };
376 mac_addresses.push(format!("{interface}:{mac}"));
377 }
378 normalize_and_validate_mac_entries(&mut mac_addresses);
379 Ok(mac_addresses)
380}
381
382fn get_machine_mac_addresses_from_ifconfig() -> io::Result<Vec<String>> {
383 let output = Command::new("ifconfig").output()?;
384 if !output.status.success() {
385 return Err(io::Error::other("`ifconfig` returned non-zero status"));
386 }
387 let mut mac_addresses = Vec::new();
388 let mut current_interface: Option<String> = None;
389 for line in String::from_utf8_lossy(&output.stdout).lines() {
390 if !line.starts_with('\t') && !line.starts_with(' ') {
391 current_interface = parse_interface_name_from_ifconfig(line);
392 continue;
393 }
394 let line = line.trim_start();
395 if !line.starts_with("ether ") {
396 continue;
397 }
398 let Some(interface) = current_interface.as_ref() else {
399 continue;
400 };
401 let Some(raw_mac) = line.split_whitespace().nth(1) else {
402 continue;
403 };
404 let Some(mac) = normalize_mac_address(raw_mac) else {
405 continue;
406 };
407 mac_addresses.push(format!("{interface}:{mac}"));
408 }
409 normalize_and_validate_mac_entries(&mut mac_addresses);
410 Ok(mac_addresses)
411}
412
413fn normalize_and_validate_mac_entries(mac_addresses: &mut Vec<String>) {
414 mac_addresses.sort();
415 mac_addresses.dedup();
416}
417
418fn parse_interface_name_from_ip_link(line: &str) -> Option<String> {
419 let mut parts = line.splitn(3, ':');
420 let _ = parts.next()?;
421 let name = parts.next()?.trim();
422 let name = name.split('@').next()?.trim();
423 if name.is_empty() || name == "lo" {
424 return None;
425 }
426 Some(name.to_string())
427}
428
429fn parse_interface_name_from_ifconfig(line: &str) -> Option<String> {
430 let name = line.split(':').next()?.trim();
431 if name.is_empty() || name == "lo" || name == "lo0" {
432 return None;
433 }
434 Some(name.to_string())
435}
436
437fn normalize_mac_address(mac: &str) -> Option<String> {
438 let mac = mac.trim().to_ascii_lowercase();
439 if mac.len() != 17 || mac == "00:00:00:00:00:00" {
440 return None;
441 }
442 for (index, ch) in mac.char_indices() {
443 if [2usize, 5, 8, 11, 14].contains(&index) {
444 if ch != ':' {
445 return None;
446 }
447 } else if !ch.is_ascii_hexdigit() {
448 return None;
449 }
450 }
451 Some(mac)
452}
453
454fn derive_msg_header_key(hostname: &str, mac_addresses: &[String]) -> String {
455 let mut normalized_mac_addresses = mac_addresses
456 .iter()
457 .map(|address| address.trim().to_ascii_lowercase())
458 .collect::<Vec<_>>();
459 normalized_mac_addresses.sort();
460 normalized_mac_addresses.dedup();
461
462 let seed = format!(
463 "{DERIVE_MSG_HEADER_KEY_TAG}|{}|{}",
464 hostname.trim().to_ascii_lowercase(),
465 normalized_mac_addresses.join("|")
466 );
467
468 let digest = digest(&SHA256, seed.as_bytes());
470
471 digest
474 .as_ref()
475 .iter()
476 .map(|byte| {
477 DERIVE_MSG_HEADER_KEY_CHARSET[(*byte as usize) % DERIVE_MSG_HEADER_KEY_CHARSET.len()]
478 as char
479 })
480 .collect()
481}
482
483fn write_machine_msg_header_key(key: &str) -> io::Result<()> {
484 let path = Path::new(MACHINE_MSG_HEADER_KEY_PATH);
485 if let Some(parent) = path.parent() {
486 std::fs::create_dir_all(parent).map_err(|error| {
487 io::Error::new(
488 error.kind(),
489 format!("failed to create directory `{}`: {error}", parent.display()),
490 )
491 })?;
492 }
493 std::fs::write(path, format!("{key}\n")).map_err(|error| {
494 io::Error::new(
495 error.kind(),
496 format!("failed to write key file `{}`: {error}", path.display()),
497 )
498 })?;
499 #[cfg(unix)]
500 std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
501 Ok(())
502}
503
504fn gen_checksum_by_key(key: &[u8]) -> ChecksumType {
505 key.iter().fold(0u32, |hash, &byte| {
506 hash.wrapping_mul(31).wrapping_add(byte as u32)
507 })
508}
509
510#[inline]
511pub fn get_checksum_for_key(datalen: DataLenType, key: &[u8]) -> ChecksumType {
513 datalen ^ gen_checksum_by_key(key)
514}
515
516pub fn process_checksum_is_ready() -> bool {
518 get_process_credential().is_ok()
519}
520
521#[inline]
522pub fn get_checksum(datalen: DataLenType) -> ChecksumType {
524 datalen ^ MSG_HEADER_KEY_STATE.hash.load(Ordering::Acquire)
525}
526
527#[inline]
528pub fn valid_checksum_for_key(datalen: DataLenType, checksum: ChecksumType, key: &[u8]) -> bool {
530 checksum == get_checksum_for_key(datalen, key)
531}
532
533#[inline]
534pub fn valid_checksum(datalen: DataLenType, checksum: ChecksumType) -> bool {
539 process_checksum_is_ready()
540 && datalen == (checksum ^ MSG_HEADER_KEY_STATE.hash.load(Ordering::Acquire))
541}
542
543pub type AesKeyType = [u8; 32];
544
545pub fn gen_random_key() -> [u8; 32] {
550 const CHARSET: &[u8] = b"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ!\"#$%&'()*+,-./:;<=>?@[\\]^_`{|}~";
551
552 let mut rng = rand::rng();
553 let mut random_key: AesKeyType = [0; 32];
554 (0..32).for_each(|i| {
555 let idx = rng.random_range(0..CHARSET.len());
556 random_key[i] = CHARSET[idx];
557 });
558
559 random_key
560}
561
562#[cfg(test)]
564mod tests {
565 #[test]
566 fn credential_debug_never_prints_key_bytes() {
567 let key = *b"0123456789abcdefghijklmnopqrstuv";
568 assert_eq!(
569 format!("{:?}", super::Credential::Admin(key)),
570 "Admin(\"[redacted]\")"
571 );
572 assert_eq!(
573 format!("{:?}", super::Credential::Temporary { key_id: 7, key }),
574 "Temporary { key_id: 7, key: \"[redacted]\" }"
575 );
576 }
577
578 #[test]
579 fn test_random_checksum() {
580 use super::*;
581 println!(
582 "{}",
583 gen_checksum_by_key(b"0123456789abcdefghijklmnopqrstuv")
584 );
585 }
586
587 #[test]
588 fn checksum_for_an_explicit_key_is_independent_of_process_state() {
589 use super::*;
590 let key = b"0123456789abcdefghijklmnopqrstuv";
591 let datalen = 32;
592 let checksum = get_checksum_for_key(datalen, key);
593 assert!(valid_checksum_for_key(datalen, checksum, key));
594 assert!(!valid_checksum_for_key(
595 datalen,
596 checksum,
597 b"abcdefghijklmnopqrstuvwxyz012345"
598 ));
599 }
600
601 #[test]
602 fn env_safe_admin_key_rejects_nul_and_accepts_printable_ascii() {
603 use super::*;
604 assert!(is_env_safe_admin_key(b"0123456789abcdefghijklmnopqrstuv"));
605 let mut with_nul = *b"0123456789abcdefghijklmnopqrstuv";
606 with_nul[8] = 0;
607 assert!(!is_env_safe_admin_key(&with_nul));
608 assert!(!is_env_safe_admin_key(b"short"));
609 }
610
611 #[tokio::test]
612 async fn clearing_the_process_credential_fails_closed_for_checksums() {
613 use super::*;
614 use crate::test_support::PROCESS_CREDENTIAL_TEST_LOCK;
615
616 let _guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await;
617 set_process_msg_header_key(Some("0123456789abcdefghijklmnopqrstuv")).unwrap();
618 let checksum = get_checksum(32);
619 assert!(valid_checksum(32, checksum));
620 set_process_msg_header_key(None).unwrap();
621 assert!(!valid_checksum(32, checksum));
622 assert!(!valid_checksum(32, 32));
623 assert!(get_process_credential().is_err());
624 }
625
626 #[test]
627 fn administrator_credentials_reject_nul_and_whitespace() {
628 use super::*;
629 let mut with_nul = *b"0123456789abcdefghijklmnopqrstuv";
630 with_nul[8] = 0;
631 assert!(parse_credential(std::str::from_utf8(&with_nul).unwrap()).is_err());
632 assert!(parse_credential("0123456789abcdefghijklmnopq rstuv").is_err());
633 }
634
635 #[test]
636 fn test_derive_msg_header_key_is_stable() {
637 use super::*;
638 let mac_addresses = vec![
639 "eth0:52:54:00:12:34:56".to_string(),
640 "ens3:02:42:ac:11:00:02".to_string(),
641 ];
642 let key1 = derive_msg_header_key("DemoHost", &mac_addresses);
643 let key2 = derive_msg_header_key("demohost", &mac_addresses);
644 assert_eq!(key1, key2);
645 assert_eq!(key1.len(), 32);
646 assert!(key1.chars().all(|ch| ch.is_ascii_alphanumeric()));
647 }
648
649 #[test]
650 fn temporary_credential_round_trip_and_checksum() {
651 use super::*;
652
653 let key = [7_u8; 32];
654 let encoded = encode_temporary_credential(0x0000_0007_0000_002a, &key);
655 assert_eq!(
656 parse_credential(&encoded).unwrap(),
657 Credential::Temporary {
658 key_id: 0x0000_0007_0000_002a,
659 key
660 }
661 );
662
663 let mut corrupted = encoded.into_bytes();
664 let last = corrupted.last_mut().unwrap();
665 *last = if *last == b'A' { b'B' } else { b'A' };
666 assert!(parse_credential(std::str::from_utf8(&corrupted).unwrap()).is_err());
667 }
668}