use base64::Engine; use serde::{Deserialize, Serialize}; use std::io; use std::path::{Path, PathBuf}; /// `ca_pem` is the trust anchor to pin, when the link carried one (the /// `ca` parameter, `wg_app_link::enroll::ca_param`). It is optional /// for compatibility with older links. Current clients require it before /// opening a transport. It is a public certificate, not a secret. #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct EnrolledServer { pub host: String, pub port: u16, pub token: String, #[serde(default)] pub ca_pem: Option, } impl EnrolledServer { /// `ca` is base64url of the certificate's DER and is rebuilt into PEM /// here, because that is what every consumer of it wants /// (`UreqTransport::new`, and the file a person points `curl --cacert` /// at). A `ca` that does not decode fails the whole link rather than /// enrolling a server with no trust anchor: the link said which /// certificate to pin, and quietly not pinning it is the one outcome /// nothing downstream could notice. pub fn parse_link(link: &str) -> Result { let query = link.split_once('?').map(|(_, q)| q).ok_or_else(|| { format!( "'{link}' has no query string (expected \ aiapp://enroll?host=...&port=...&token=...)" ) })?; let mut host = None; let mut port = None; let mut token = None; let mut ca = None; for pair in query.split('&') { let Some((key, value)) = pair.split_once('=') else { continue; }; let value = percent_decode(value); match key { "host" => host = Some(value), "port" => port = Some(value), "token" => token = Some(value), "ca" => ca = Some(value), _ => {} } } let host = host.ok_or_else(|| format!("'{link}' is missing 'host'"))?; let port_str = port.ok_or_else(|| format!("'{link}' is missing 'port'"))?; let port: u16 = port_str .parse() .map_err(|e| format!("'{link}''s port ('{port_str}') is not a number: {e}"))?; let token = token.ok_or_else(|| format!("'{link}' is missing 'token'"))?; let ca_pem = ca.map(|ca| pem_from_link_param(&ca)).transpose()?; Ok(Self { host, port, token, ca_pem, }) } pub fn base_url(&self) -> String { format!("https://{}:{}", self.host, self.port) } } fn pem_from_link_param(ca: &str) -> Result { let der = base64::engine::general_purpose::URL_SAFE_NO_PAD .decode(ca.as_bytes()) .map_err(|e| format!("the link's 'ca' is not base64url ({e})"))?; let body = base64::engine::general_purpose::STANDARD.encode(&der); let mut pem = String::from("-----BEGIN CERTIFICATE-----\n"); for line in body.as_bytes().chunks(64) { pem.push_str(std::str::from_utf8(line).expect("base64 is ASCII")); pem.push('\n'); } pem.push_str("-----END CERTIFICATE-----\n"); Ok(pem) } /// Where one client keeps the enrollment it should not have to be told /// about a second time. `dir` is the caller's, because that is the only /// part that differs by platform -- see this module's doc. pub struct EnrollmentStore { dir: PathBuf, } impl EnrollmentStore { pub fn new(dir: impl Into) -> Self { Self { dir: dir.into() } } pub fn dir(&self) -> &Path { &self.dir } fn file(&self) -> PathBuf { self.dir.join("enrollment.json") } pub fn save(&self, server: &EnrolledServer) -> io::Result<()> { std::fs::create_dir_all(&self.dir)?; let path = self.file(); let json = serde_json::to_vec_pretty(server) .expect("EnrolledServer holds nothing that fails to serialise"); std::fs::write(&path, json)?; #[cfg(unix)] { use std::os::unix::fs::PermissionsExt; std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o600))?; } Ok(()) } pub fn load(&self) -> io::Result> { let path = self.file(); match std::fs::read(&path) { Ok(bytes) => { let server = serde_json::from_slice(&bytes).map_err(|e| { io::Error::new( io::ErrorKind::InvalidData, format!("{} is not a valid enrollment ({e})", path.display()), ) })?; Ok(Some(server)) } Err(e) if e.kind() == io::ErrorKind::NotFound => Ok(None), Err(e) => Err(e), } } } fn percent_decode(s: &str) -> String { let bytes = s.as_bytes(); let mut out = Vec::with_capacity(bytes.len()); let mut i = 0; while i < bytes.len() { if bytes[i] == b'%' && i + 2 < bytes.len() && let Ok(byte) = u8::from_str_radix(std::str::from_utf8(&bytes[i + 1..i + 3]).unwrap_or(""), 16) { out.push(byte); i += 3; continue; } out.push(bytes[i]); i += 1; } String::from_utf8_lossy(&out).into_owned() } #[cfg(test)] mod tests { use super::*; #[test] fn parses_host_port_and_token() { let server = EnrolledServer::parse_link("aiapp://enroll?host=127.0.0.1&port=8547&token=abcDEF123") .unwrap(); assert_eq!( server, EnrolledServer { host: "127.0.0.1".to_string(), port: 8547, token: "abcDEF123".to_string(), ca_pem: None, } ); assert_eq!(server.base_url(), "https://127.0.0.1:8547"); } #[test] fn field_order_does_not_matter() { let server = EnrolledServer::parse_link("aiapp://enroll?token=tok&port=443&host=example.com") .unwrap(); assert_eq!(server.host, "example.com"); assert_eq!(server.port, 443); assert_eq!(server.token, "tok"); } #[test] fn a_percent_encoded_token_is_decoded() { let server = EnrolledServer::parse_link("aiapp://enroll?host=h&port=1&token=a%2Bb%2Fc").unwrap(); assert_eq!(server.token, "a+b/c"); } #[test] fn a_missing_field_is_named_in_the_error() { let err = EnrolledServer::parse_link("aiapp://enroll?host=h&port=1").unwrap_err(); assert!( err.contains("token"), "error should name the missing field: {err}" ); } #[test] fn a_ca_in_the_link_comes_back_as_pem() { let der = [0x30u8, 0x82, 0x01, 0xfb, 0x3e, 0x7f]; let param = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(der); let server = EnrolledServer::parse_link(&format!("aiapp://enroll?host=h&port=1&token=t&ca={param}")) .unwrap(); let pem = server.ca_pem.expect("the link carried a CA"); assert!(pem.starts_with("-----BEGIN CERTIFICATE-----\n"), "{pem}"); assert!( pem.trim_end().ends_with("-----END CERTIFICATE-----"), "{pem}" ); assert_eq!( base64::engine::general_purpose::STANDARD .decode( pem.lines() .filter(|l| !l.starts_with("-----")) .collect::() ) .unwrap(), der ); } #[test] fn no_ca_parameter_is_none_not_an_error() { let server = EnrolledServer::parse_link("aiapp://enroll?host=h&port=1&token=t").unwrap(); assert_eq!(server.ca_pem, None); } #[test] fn a_ca_that_does_not_decode_fails_the_link() { let err = EnrolledServer::parse_link("aiapp://enroll?host=h&port=1&token=t&ca=not!base64url") .unwrap_err(); assert!(err.contains("ca"), "{err}"); } #[test] fn a_saved_enrollment_reads_back_the_same() { let dir = tempfile::tempdir().unwrap(); let store = EnrollmentStore::new(dir.path()); let server = EnrolledServer { host: "127.0.0.1".to_string(), port: 8547, token: "tok".to_string(), ca_pem: Some("-----BEGIN CERTIFICATE-----\nQUJD\n-----END CERTIFICATE-----\n".into()), }; store.save(&server).unwrap(); assert_eq!(store.load().unwrap(), Some(server)); } #[test] fn nothing_saved_yet_is_none_not_an_error() { let dir = tempfile::tempdir().unwrap(); assert_eq!(EnrollmentStore::new(dir.path()).load().unwrap(), None); } #[test] fn an_enrollment_without_a_ca_still_loads() { let dir = tempfile::tempdir().unwrap(); let store = EnrollmentStore::new(dir.path()); std::fs::create_dir_all(dir.path()).unwrap(); std::fs::write( dir.path().join("enrollment.json"), br#"{"host":"h","port":1,"token":"t"}"#, ) .unwrap(); assert_eq!(store.load().unwrap().unwrap().ca_pem, None); } #[test] #[cfg(unix)] fn the_saved_file_is_owner_only() { use std::os::unix::fs::PermissionsExt; let dir = tempfile::tempdir().unwrap(); let store = EnrollmentStore::new(dir.path()); store .save(&EnrolledServer { host: "h".to_string(), port: 1, token: "t".to_string(), ca_pem: None, }) .unwrap(); let mode = std::fs::metadata(dir.path().join("enrollment.json")) .unwrap() .permissions() .mode(); assert_eq!(mode & 0o777, 0o600); } #[test] fn a_corrupt_file_is_named_in_the_error() { let dir = tempfile::tempdir().unwrap(); std::fs::write(dir.path().join("enrollment.json"), b"not json").unwrap(); let err = EnrollmentStore::new(dir.path()).load().unwrap_err(); assert!(err.to_string().contains("enrollment.json")); } #[test] fn a_non_numeric_port_is_named_in_the_error() { let err = EnrolledServer::parse_link("aiapp://enroll?host=h&port=x&token=t").unwrap_err(); assert!( err.contains("port"), "error should name the offending field: {err}" ); } }