use std::fmt::Write as _; use std::fs; use std::io::{BufRead, BufReader, Write}; use std::os::unix::net::UnixStream; use std::path::{Path, PathBuf}; use anyhow::{Context, Result, bail}; #[derive(Clone)] pub struct Refresh { socket: PathBuf, token: Vec, } impl Refresh { pub fn from_files( socket: Option, token_file: Option, ) -> Result> { let (Some(socket), Some(token_file)) = (socket, token_file) else { return Ok(None); }; let token = fs::read(&token_file) .with_context(|| format!("reading refresh token {}", token_file.display()))?; let mut token = token.as_slice(); if let Some(without_lf) = token.strip_suffix(b"\n") { token = without_lf; } if let Some(without_cr) = token.strip_suffix(b"\r") { token = without_cr; } if token.is_empty() || !token.iter().all(u8::is_ascii_graphic) { bail!("refresh token must contain visible ASCII without whitespace"); } Ok(Some(Self { socket, token: token.to_vec(), })) } pub fn send(&self, user: &str, repo: &str) -> Result<()> { refresh(&self.socket, &self.token, user, repo) } } fn refresh(socket: &Path, token: &[u8], user: &str, repo: &str) -> Result<()> { let mut stream = UnixStream::connect(socket) .with_context(|| format!("connecting to sorcery socket {}", socket.display()))?; write!( stream, "POST /-/refresh/{}/{} HTTP/1.1\r\nHost: sorcery\r\nAuthorization: Bearer ", encode_component(user), encode_component(repo), )?; stream.write_all(token)?; stream.write_all(b"\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")?; stream.flush()?; let mut status = String::new(); BufReader::new(stream).read_line(&mut status)?; let code = status.split_ascii_whitespace().nth(1); if code != Some("200") { bail!("sorcery refresh returned {}", status.trim_end()); } Ok(()) } fn encode_component(component: &str) -> String { let mut encoded = String::with_capacity(component.len()); for byte in component.bytes() { if byte.is_ascii_alphanumeric() || b"-._~".contains(&byte) { encoded.push(char::from(byte)); } else { write!(encoded, "%{byte:02X}").expect("writing to a String cannot fail"); } } encoded } #[cfg(test)] mod tests { use std::os::unix::net::UnixListener; use std::thread; use super::*; #[test] fn sends_authenticated_refresh_request() { let temp = tempfile::tempdir().unwrap(); let socket = temp.path().join("socket"); let listener = UnixListener::bind(&socket).unwrap(); let server = thread::spawn(move || { let (mut connection, _) = listener.accept().unwrap(); let mut request = String::new(); let mut reader = BufReader::new(connection.try_clone().unwrap()); loop { let mut line = String::new(); reader.read_line(&mut line).unwrap(); request.push_str(&line); if line == "\r\n" { break; } } connection .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\nrefreshed\n") .unwrap(); request }); refresh(&socket, b"secret", "an user", "a/repo").unwrap(); let request = server.join().unwrap(); assert!(request.starts_with("POST /-/refresh/an%20user/a%2Frepo HTTP/1.1\r\n")); assert!(request.contains("Authorization: Bearer secret\r\n")); } #[test] fn percent_encodes_path_components() { assert_eq!(encode_component("hello world/ΓΈ"), "hello%20world%2F%C3%B8"); } #[test] fn incomplete_configuration_disables_refresh() { assert!(Refresh::from_files(None, None).unwrap().is_none()); assert!( Refresh::from_files(Some("/tmp/sock".into()), None) .unwrap() .is_none() ); } #[test] fn rejects_whitespace_in_token() { let temp = tempfile::NamedTempFile::new().unwrap(); fs::write(temp.path(), "not valid\ninside").unwrap(); let error = Refresh::from_files(Some("/tmp/sock".into()), Some(temp.path().into())); assert!(error.is_err()); } }