use std::io::{self, Read, Write};
use std::os::unix::fs::{DirBuilderExt, PermissionsExt};
use std::os::unix::net::{UnixListener, UnixStream};
use std::path::{Path, PathBuf};
use std::time::Duration;
use teksilo_automation::wire::Endpoint;
use super::{BoundTransport, TransportListener, TransportStream};
const MAX_SOCKET_PATH: usize = 100;
struct SocketStream(UnixStream);
impl Read for SocketStream {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
self.0.read(buf).map_err(timed_out_if_would_block)
}
}
impl Write for SocketStream {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.0.write(buf)
}
fn flush(&mut self) -> io::Result<()> {
self.0.flush()
}
}
impl TransportStream for SocketStream {
fn set_read_timeout(&self, timeout: Option<Duration>) -> io::Result<()> {
self.0.set_read_timeout(timeout)
}
}
fn timed_out_if_would_block(e: io::Error) -> io::Error {
if e.kind() == io::ErrorKind::WouldBlock {
io::Error::new(io::ErrorKind::TimedOut, "the socket read deadline expired")
} else {
e
}
}
struct SocketListener {
listener: UnixListener,
dir: PathBuf,
}
impl TransportListener for SocketListener {
fn accept(&mut self) -> io::Result<Box<dyn TransportStream>> {
let (stream, _addr) = self.listener.accept()?;
Ok(Box::new(SocketStream(stream)))
}
}
impl Drop for SocketListener {
fn drop(&mut self) {
let _ = std::fs::remove_dir_all(&self.dir);
}
}
fn candidate_dirs(pid: u32) -> Vec<PathBuf> {
let mut dirs = vec![teksilo_automation::wire::EndpointFile::dir().join(format!("{pid}.d"))];
dirs.push(PathBuf::from(format!("/tmp/tka-{pid}")));
dirs
}
pub(super) fn bind(pid: u32) -> io::Result<BoundTransport> {
bind_over(&candidate_dirs(pid))
}
fn bind_over(candidates: &[PathBuf]) -> io::Result<BoundTransport> {
let mut last_err = None;
for dir in candidates {
let path = dir.join("s");
if path.as_os_str().len() > MAX_SOCKET_PATH {
last_err = Some(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"socket path {} is {} bytes, over the {MAX_SOCKET_PATH}-byte limit",
path.display(),
path.as_os_str().len()
),
));
continue;
}
match bind_at(dir, &path) {
Ok(listener) => {
let address = path.to_string_lossy().into_owned();
return Ok(BoundTransport {
listener: Box::new(listener),
endpoint: Endpoint::unix(address),
});
}
Err(e) => last_err = Some(e),
}
}
Err(last_err.unwrap_or_else(|| {
io::Error::other("no usable directory for the automation bridge socket")
}))
}
fn bind_at(dir: &Path, path: &Path) -> io::Result<SocketListener> {
let _ = std::fs::remove_dir_all(dir);
if let Some(parent) = dir.parent() {
teksilo_automation::wire::create_private_dir(parent)?;
}
std::fs::DirBuilder::new().mode(0o700).create(dir)?;
let _ = std::fs::remove_file(path);
let listener = UnixListener::bind(path)?;
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
Ok(SocketListener {
listener,
dir: dir.to_path_buf(),
})
}
pub(super) fn connect(address: &str) -> io::Result<Box<dyn TransportStream>> {
Ok(Box::new(SocketStream(UnixStream::connect(address)?)))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn socket_is_owner_only_in_an_owner_only_directory() {
let pid = std::process::id().wrapping_add(7);
let bound = match bind(pid) {
Ok(b) => b,
Err(e) if e.kind() == io::ErrorKind::PermissionDenied => return,
Err(e) => panic!("bind failed: {e}"),
};
let path = PathBuf::from(&bound.endpoint.address);
let sock_mode = std::fs::metadata(&path).unwrap().permissions().mode() & 0o777;
let dir_mode = std::fs::metadata(path.parent().unwrap())
.unwrap()
.permissions()
.mode()
& 0o777;
assert_eq!(sock_mode, 0o600, "socket must be owner-only");
assert_eq!(dir_mode, 0o700, "directory must be owner-only");
}
#[test]
fn dropping_the_listener_removes_the_directory() {
let pid = std::process::id().wrapping_add(8);
let bound = match bind(pid) {
Ok(b) => b,
Err(e) if e.kind() == io::ErrorKind::PermissionDenied => return,
Err(e) => panic!("bind failed: {e}"),
};
let path = PathBuf::from(&bound.endpoint.address);
let dir = path.parent().unwrap().to_path_buf();
assert!(dir.exists());
drop(bound);
assert!(
!dir.exists(),
"the per-process directory must not outlive the bridge"
);
}
#[test]
fn an_overlong_preferred_path_falls_back_instead_of_failing() {
let deep = PathBuf::from("/tmp/".to_string() + &"d".repeat(120));
assert!(
deep.join("s").as_os_str().len() > MAX_SOCKET_PATH,
"the simulated preferred directory must actually overflow"
);
let fallback = candidate_dirs(1234).pop().expect("a fallback candidate");
assert!(
fallback.join("s").as_os_str().len() <= MAX_SOCKET_PATH,
"the fallback must always fit"
);
let pid = std::process::id().wrapping_add(9);
let candidates = vec![deep.clone(), PathBuf::from(format!("/tmp/tka-{pid}"))];
let bound = match bind_over(&candidates) {
Ok(b) => b,
Err(e) if e.kind() == io::ErrorKind::PermissionDenied => return,
Err(e) => panic!("bind should have fallen back, but failed: {e}"),
};
assert!(
!bound.endpoint.address.starts_with(deep.to_str().unwrap()),
"the overlong preferred path must be skipped, got {}",
bound.endpoint.address
);
}
}