char/sorcery

static-files based git repo viewer

git clone https://git.t4t.associates/char/sorcery

Charlotte Somadd sorcery-ssh with sorcery-ssh-tui8f07a0e

main
4.4 KiB138 linesraw
1use std::fmt::Write as _;
2use std::fs;
3use std::io::{BufRead, BufReader, Write};
4use std::os::unix::net::UnixStream;
5use std::path::{Path, PathBuf};
6
7use anyhow::{Context, Result, bail};
8
9#[derive(Clone)]
10pub struct Refresh {
11    socket: PathBuf,
12    token: Vec<u8>,
13}
14
15impl Refresh {
16    pub fn from_files(
17        socket: Option<PathBuf>,
18        token_file: Option<PathBuf>,
19    ) -> Result<Option<Self>> {
20        let (Some(socket), Some(token_file)) = (socket, token_file) else {
21            return Ok(None);
22        };
23        let token = fs::read(&token_file)
24            .with_context(|| format!("reading refresh token {}", token_file.display()))?;
25        let mut token = token.as_slice();
26        if let Some(without_lf) = token.strip_suffix(b"\n") {
27            token = without_lf;
28        }
29        if let Some(without_cr) = token.strip_suffix(b"\r") {
30            token = without_cr;
31        }
32        if token.is_empty() || !token.iter().all(u8::is_ascii_graphic) {
33            bail!("refresh token must contain visible ASCII without whitespace");
34        }
35        Ok(Some(Self {
36            socket,
37            token: token.to_vec(),
38        }))
39    }
40
41    pub fn send(&self, user: &str, repo: &str) -> Result<()> {
42        refresh(&self.socket, &self.token, user, repo)
43    }
44}
45
46fn refresh(socket: &Path, token: &[u8], user: &str, repo: &str) -> Result<()> {
47    let mut stream = UnixStream::connect(socket)
48        .with_context(|| format!("connecting to sorcery socket {}", socket.display()))?;
49    write!(
50        stream,
51        "POST /-/refresh/{}/{} HTTP/1.1\r\nHost: sorcery\r\nAuthorization: Bearer ",
52        encode_component(user),
53        encode_component(repo),
54    )?;
55    stream.write_all(token)?;
56    stream.write_all(b"\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")?;
57    stream.flush()?;
58
59    let mut status = String::new();
60    BufReader::new(stream).read_line(&mut status)?;
61    let code = status.split_ascii_whitespace().nth(1);
62    if code != Some("200") {
63        bail!("sorcery refresh returned {}", status.trim_end());
64    }
65    Ok(())
66}
67
68fn encode_component(component: &str) -> String {
69    let mut encoded = String::with_capacity(component.len());
70    for byte in component.bytes() {
71        if byte.is_ascii_alphanumeric() || b"-._~".contains(&byte) {
72            encoded.push(char::from(byte));
73        } else {
74            write!(encoded, "%{byte:02X}").expect("writing to a String cannot fail");
75        }
76    }
77    encoded
78}
79
80#[cfg(test)]
81mod tests {
82    use std::os::unix::net::UnixListener;
83    use std::thread;
84
85    use super::*;
86
87    #[test]
88    fn sends_authenticated_refresh_request() {
89        let temp = tempfile::tempdir().unwrap();
90        let socket = temp.path().join("socket");
91        let listener = UnixListener::bind(&socket).unwrap();
92        let server = thread::spawn(move || {
93            let (mut connection, _) = listener.accept().unwrap();
94            let mut request = String::new();
95            let mut reader = BufReader::new(connection.try_clone().unwrap());
96            loop {
97                let mut line = String::new();
98                reader.read_line(&mut line).unwrap();
99                request.push_str(&line);
100                if line == "\r\n" {
101                    break;
102                }
103            }
104            connection
105                .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\nrefreshed\n")
106                .unwrap();
107            request
108        });
109
110        refresh(&socket, b"secret", "an user", "a/repo").unwrap();
111        let request = server.join().unwrap();
112        assert!(request.starts_with("POST /-/refresh/an%20user/a%2Frepo HTTP/1.1\r\n"));
113        assert!(request.contains("Authorization: Bearer secret\r\n"));
114    }
115
116    #[test]
117    fn percent_encodes_path_components() {
118        assert_eq!(encode_component("hello world/ø"), "hello%20world%2F%C3%B8");
119    }
120
121    #[test]
122    fn incomplete_configuration_disables_refresh() {
123        assert!(Refresh::from_files(None, None).unwrap().is_none());
124        assert!(
125            Refresh::from_files(Some("/tmp/sock".into()), None)
126                .unwrap()
127                .is_none()
128        );
129    }
130
131    #[test]
132    fn rejects_whitespace_in_token() {
133        let temp = tempfile::NamedTempFile::new().unwrap();
134        fs::write(temp.path(), "not valid\ninside").unwrap();
135        let error = Refresh::from_files(Some("/tmp/sock".into()), Some(temp.path().into()));
136        assert!(error.is_err());
137    }
138}