use std::io::{BufRead, BufReader, Write};
use std::os::unix::net::{UnixListener, UnixStream};
use std::path::{Path, PathBuf};
use crate::cmd::Command;
const SOCKET_NAME: &str = "tree-space.sock";
pub fn socket_path() -> PathBuf {
let dir = std::env::var_os("XDG_RUNTIME_DIR")
.map(PathBuf::from)
.filter(|p| !p.as_os_str().is_empty())
.unwrap_or_else(|| PathBuf::from("/tmp"));
dir.join(SOCKET_NAME)
}
pub fn deliver(cmd: &Command) -> bool {
deliver_at(&socket_path(), cmd)
}
pub fn deliver_at(path: &Path, cmd: &Command) -> bool {
let Ok(mut stream) = UnixStream::connect(path) else {
return false;
};
let line = format!("{}\n", cmd.encode());
if stream.write_all(line.as_bytes()).is_err() || stream.flush().is_err() {
return false;
}
true
}
pub fn bind() -> std::io::Result<UnixListener> {
bind_at(&socket_path())
}
pub fn bind_at(path: &Path) -> std::io::Result<UnixListener> {
if let Some(parent) = path.parent() {
let _ = std::fs::create_dir_all(parent);
}
match UnixListener::bind(path) {
Ok(listener) => Ok(listener),
Err(err) if err.kind() == std::io::ErrorKind::AddrInUse => {
if UnixStream::connect(path).is_ok() {
return Err(err);
}
let _ = std::fs::remove_file(path);
UnixListener::bind(path)
}
Err(err) => Err(err),
}
}
pub fn spawn_listener(listener: UnixListener, deliver: impl Fn(Command) + Send + 'static) {
std::thread::spawn(move || {
for connection in listener.incoming() {
let Ok(stream) = connection else { continue };
let mut reader = BufReader::new(stream);
let mut line = String::new();
if reader.read_line(&mut line).is_err() {
continue;
}
if let Some(cmd) = Command::decode(line.trim_end()) {
deliver(cmd);
}
}
});
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::mpsc;
use std::time::Duration;
#[test]
fn deliver_and_listener_round_trip() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(SOCKET_NAME);
let listener = bind_at(&path).unwrap();
let (tx, rx) = mpsc::channel();
spawn_listener(listener, move |cmd| tx.send(cmd).unwrap());
let expected = Command {
side: Some(crate::config::PanelSide::Right),
roots: vec![PathBuf::from("/home/eolu/x")],
..Command::default()
};
assert!(deliver_at(&path, &expected));
assert_eq!(rx.recv_timeout(Duration::from_secs(2)), Ok(expected));
}
#[test]
fn deliver_fails_without_a_server() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(SOCKET_NAME);
assert!(!deliver_at(&path, &Command::default()));
}
#[test]
fn deliver_toggle_round_trip() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(SOCKET_NAME);
let listener = bind_at(&path).unwrap();
let (tx, rx) = mpsc::channel();
spawn_listener(listener, move |cmd| tx.send(cmd).unwrap());
assert!(deliver_at(&path, &Command::default()));
assert_eq!(rx.recv_timeout(Duration::from_secs(2)), Ok(Command::default()));
}
#[test]
fn bind_replaces_a_stale_socket() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(SOCKET_NAME);
std::fs::write(&path, b"stale").unwrap();
let listener = bind_at(&path).unwrap();
assert!(listener.local_addr().is_ok());
}
#[test]
fn bind_refuses_when_a_live_server_holds_the_socket() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(SOCKET_NAME);
let _server = bind_at(&path).unwrap();
assert!(bind_at(&path).is_err());
}
}