mod connect;
use super::PersistentService;
use anyhow::{Context, Result, bail, ensure};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use std::{
fs::{self, DirBuilder, File, OpenOptions},
io::{Read, Write},
os::{
fd::AsRawFd,
unix::{
fs::{DirBuilderExt, FileTypeExt, MetadataExt, OpenOptionsExt, PermissionsExt},
net::{UnixListener, UnixStream},
},
},
path::{Path, PathBuf},
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
thread,
time::{Duration, Instant},
};
const VERSION: &str = env!("CARGO_PKG_VERSION");
const HANDSHAKE_LIMIT: usize = 4096;
const CLIENT_LIMIT: usize = 32;
const TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Clone, Serialize, Deserialize)]
pub(crate) struct Identity {
pub workspace: PathBuf,
pub state_root: PathBuf,
pub socket: PathBuf,
}
impl Identity {
pub(crate) fn resolve(workspace: &Path, root: &Path) -> Result<Self> {
let workspace = workspace.canonicalize().context("workspace must exist")?;
ensure!(workspace.is_dir(), "workspace must be a directory");
fs::create_dir_all(root)?;
let state_root = root.canonicalize()?;
let mut hash = Sha256::new();
use std::os::unix::ffi::OsStrExt;
for path in [&workspace, &state_root] {
let bytes = path.as_os_str().as_bytes();
hash.update((bytes.len() as u64).to_le_bytes());
hash.update(bytes);
}
let key: String = hash
.finalize()
.iter()
.map(|byte| format!("{byte:02x}"))
.collect();
let directory = PathBuf::from(format!("/tmp/magi-{}", current_uid()));
match DirBuilder::new().mode(0o700).create(&directory) {
Ok(()) => {}
Err(error) if error.kind() == std::io::ErrorKind::AlreadyExists => {}
Err(error) => return Err(error.into()),
}
check_path(&directory, Kind::Directory)?;
Ok(Self {
workspace,
state_root,
socket: directory.join(format!("{}.sock", &key[..48])),
})
}
}
#[derive(Clone, Copy)]
enum Kind {
Directory,
Socket,
Lock,
}
fn current_uid() -> u32 {
unsafe { libc::geteuid() }
}
fn check_path(path: &Path, kind: Kind) -> Result<()> {
let metadata = fs::symlink_metadata(path)?;
let valid_type = match kind {
Kind::Directory => metadata.is_dir(),
Kind::Socket => metadata.file_type().is_socket(),
Kind::Lock => metadata.is_file(),
};
ensure!(
valid_type
&& !metadata.file_type().is_symlink()
&& metadata.uid() == current_uid()
&& metadata.mode() & 0o077 == 0,
"unsafe daemon endpoint"
);
Ok(())
}
fn check_peer(stream: &UnixStream) -> Result<()> {
#[cfg(target_os = "macos")]
let uid = {
let mut uid = 0;
let mut gid = 0;
ensure!(
unsafe { libc::getpeereid(stream.as_raw_fd(), &mut uid, &mut gid) } == 0,
"cannot authenticate daemon peer"
);
uid
};
#[cfg(target_os = "linux")]
let uid = {
let mut credentials: libc::ucred = unsafe { std::mem::zeroed() };
let mut length = std::mem::size_of::<libc::ucred>() as libc::socklen_t;
ensure!(
unsafe {
libc::getsockopt(
stream.as_raw_fd(),
libc::SOL_SOCKET,
libc::SO_PEERCRED,
(&mut credentials as *mut libc::ucred).cast(),
&mut length,
)
} == 0
&& length as usize == std::mem::size_of::<libc::ucred>(),
"cannot authenticate daemon peer"
);
credentials.uid
};
ensure!(uid == current_uid(), "daemon peer belongs to another user");
Ok(())
}
#[derive(Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
struct Hello {
version: String,
workspace: PathBuf,
state_root: PathBuf,
action: String,
}
#[derive(Serialize, Deserialize)]
struct Reply {
version: String,
workspace: PathBuf,
state_root: PathBuf,
status: String,
}
fn read_line(stream: &mut UnixStream, limit: usize) -> Result<Vec<u8>> {
stream.set_nonblocking(true)?;
let deadline = Instant::now() + TIMEOUT;
let mut bytes = Vec::new();
loop {
ensure!(Instant::now() < deadline, "daemon frame timeout");
let mut byte = [0];
match stream.read(&mut byte) {
Ok(0) => bail!("daemon connection closed"),
Ok(_) if byte[0] == b'\n' => return Ok(bytes),
Ok(_) => {
ensure!(bytes.len() < limit, "daemon frame too large");
bytes.push(byte[0]);
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(2))
}
Err(error) => return Err(error.into()),
}
}
}
fn send_json(stream: &mut UnixStream, value: &impl Serialize) -> Result<()> {
stream.set_nonblocking(true)?;
let mut bytes = serde_json::to_vec(value)?;
bytes.push(b'\n');
let deadline = Instant::now() + TIMEOUT;
let mut written = 0;
while written < bytes.len() {
ensure!(Instant::now() < deadline, "daemon write timeout");
match stream.write(&bytes[written..]) {
Ok(0) => bail!("daemon connection closed"),
Ok(count) => written += count,
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(2))
}
Err(error) => return Err(error.into()),
}
}
Ok(())
}
fn request(identity: &Identity, action: &str) -> Result<Reply> {
check_path(
identity
.socket
.parent()
.context("missing endpoint directory")?,
Kind::Directory,
)?;
check_path(&identity.socket, Kind::Socket)?;
let mut stream = connect::connect(&identity.socket).context("connect daemon socket")?;
check_peer(&stream)?;
send_json(
&mut stream,
&Hello {
version: VERSION.into(),
workspace: identity.workspace.clone(),
state_root: identity.state_root.clone(),
action: action.into(),
},
)?;
let reply: Reply = serde_json::from_slice(&read_line(&mut stream, HANDSHAKE_LIMIT)?)?;
ensure!(
reply.version == VERSION,
"incompatible daemon version; stop it using its installed version when idle"
);
ensure!(
reply.workspace == identity.workspace && reply.state_root == identity.state_root,
"daemon identity mismatch"
);
Ok(reply)
}
pub(crate) fn control(identity: &Identity, action: &str) -> Result<()> {
let reply = request(identity, action)?;
ensure!(
reply.status != "busy",
"daemon is busy; no work was stopped"
);
println!(
"{}",
serde_json::to_string(
&serde_json::json!({"status": reply.status, "version": reply.version, "workspace": reply.workspace, "state_root": reply.state_root, "socket": identity.socket})
)?
);
Ok(())
}
pub(crate) fn start_or_connect(identity: &Identity, executable: &Path) -> Result<()> {
match fs::symlink_metadata(&identity.socket) {
Ok(_) => {
check_path(&identity.socket, Kind::Socket)?;
match connect::connect(&identity.socket) {
Ok(_) => return control(identity, "status"),
Err(error) if error.kind() == std::io::ErrorKind::ConnectionRefused => {}
Err(error) => return Err(error.into()),
}
}
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {}
Err(error) => return Err(error.into()),
}
ensure!(
executable.is_absolute(),
"daemon executable must be absolute"
);
let metadata = fs::symlink_metadata(executable)?;
ensure!(
metadata.is_file()
&& !metadata.file_type().is_symlink()
&& metadata.mode() & 0o222 == 0
&& (metadata.uid() == current_uid() || metadata.uid() == 0),
"untrusted daemon executable"
);
for directory in executable
.parent()
.context("missing executable directory")?
.ancestors()
{
let metadata = fs::symlink_metadata(directory)?;
ensure!(
metadata.is_dir()
&& !metadata.file_type().is_symlink()
&& (metadata.uid() == current_uid() || metadata.uid() == 0)
&& (metadata.mode() & 0o022 == 0 || metadata.mode() & 0o1000 != 0),
"untrusted daemon executable directory"
);
}
let mut child = std::process::Command::new(executable)
.args(["daemon", "foreground", "--detached", "--workspace"])
.arg(&identity.workspace)
.arg("--state-root")
.arg(&identity.state_root)
.env("MC_HOME", &identity.state_root)
.current_dir(&identity.workspace)
.stdin(std::process::Stdio::null())
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.spawn()
.context("daemon launch failed")?;
let deadline = Instant::now() + Duration::from_secs(15);
loop {
if let Ok(reply) = request(identity, "status") {
ensure!(reply.status == "ready", "daemon is not ready");
return control(identity, "status");
}
if let Some(status) = child.try_wait()? {
if status.success() {
return control(identity, "status");
}
}
ensure!(
Instant::now() < deadline,
"daemon readiness timeout; no process was killed; inspect status before retrying"
);
thread::sleep(Duration::from_millis(50));
}
}
fn lock_identity(identity: &Identity) -> Result<File> {
let path = identity.socket.with_extension("lock");
let file = OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(false)
.mode(0o600)
.custom_flags(libc::O_NOFOLLOW | libc::O_CLOEXEC)
.open(&path)?;
check_path(&path, Kind::Lock)?;
ensure!(file.metadata()?.nlink() == 1, "unsafe daemon lock links");
ensure!(
unsafe { libc::flock(file.as_raw_fd(), libc::LOCK_EX | libc::LOCK_NB) } == 0,
"daemon already starting or running"
);
Ok(file)
}
pub(crate) fn foreground(identity: Identity, detached: bool) -> Result<()> {
if detached {
ensure!(
unsafe { libc::setsid() } >= 0,
"cannot detach daemon session"
);
}
ensure!(current_uid() != 0, "root daemons are unsupported");
let _lock = lock_identity(&identity)?;
if fs::symlink_metadata(&identity.socket).is_ok() {
check_path(&identity.socket, Kind::Socket)?;
match connect::connect(&identity.socket) {
Ok(_) => bail!("daemon endpoint already live"),
Err(error) if error.kind() == std::io::ErrorKind::ConnectionRefused => {
fs::remove_file(&identity.socket)?
}
Err(error) => return Err(error.into()),
}
}
let service = Arc::new(PersistentService::start_unix(
identity.workspace.clone(),
identity.state_root.clone(),
)?);
let listener = UnixListener::bind(&identity.socket)?;
fs::set_permissions(&identity.socket, fs::Permissions::from_mode(0o600))?;
listener.set_nonblocking(true)?;
let mut clients: Vec<thread::JoinHandle<()>> = Vec::new();
let stopping = Arc::new(AtomicBool::new(false));
let mut retry_delay = Duration::from_millis(10);
let result = loop {
if service.is_finished() {
break Ok(());
}
let mut index = 0;
while index < clients.len() {
if clients[index].is_finished() {
let _ = clients.swap_remove(index).join();
} else {
index += 1;
}
}
match listener.accept() {
Ok((stream, _)) if clients.len() < CLIENT_LIMIT => {
let service = Arc::clone(&service);
let identity = identity.clone();
let stopping = Arc::clone(&stopping);
match thread::Builder::new().spawn(move || {
let _ = serve(stream, &identity, &service, &stopping);
}) {
Ok(client) => {
retry_delay = Duration::from_millis(10);
clients.push(client);
}
Err(error) if recoverable_listener_error(&error) => {
backoff_listener(&mut retry_delay);
}
Err(error) => break Err(error.into()),
}
}
Ok(_) => thread::sleep(Duration::from_millis(10)),
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
retry_delay = Duration::from_millis(10);
thread::sleep(retry_delay);
}
Err(error) if recoverable_listener_error(&error) => {
backoff_listener(&mut retry_delay);
}
Err(error) => break Err(error.into()),
}
};
drop(listener);
stopping.store(true, Ordering::Release);
for client in clients {
let _ = client.join();
}
drop(service);
check_path(&identity.socket, Kind::Socket)?;
fs::remove_file(&identity.socket)?;
result
}
fn backoff_listener(delay: &mut Duration) {
thread::sleep(*delay);
*delay = (*delay * 2).min(Duration::from_millis(250));
}
fn recoverable_listener_error(error: &std::io::Error) -> bool {
matches!(
error.kind(),
std::io::ErrorKind::Interrupted | std::io::ErrorKind::ConnectionAborted
) || matches!(
error.raw_os_error(),
Some(
libc::EMFILE
| libc::ENFILE
| libc::ENOBUFS
| libc::ENOMEM
| libc::EPROTO
| libc::EAGAIN
)
)
}
fn serve(
mut stream: UnixStream,
identity: &Identity,
service: &PersistentService,
stopping: &AtomicBool,
) -> Result<()> {
check_peer(&stream)?;
let hello: Hello = serde_json::from_slice(&read_line(&mut stream, HANDSHAKE_LIMIT)?)?;
ensure!(
hello.workspace == identity.workspace && hello.state_root == identity.state_root,
"daemon identity mismatch"
);
let status = if hello.version != VERSION {
"incompatible"
} else {
match hello.action.as_str() {
"status" | "connect" => "ready",
"stop" => {
if service.stop_if_idle()? {
"stopped"
} else {
"busy"
}
}
_ => bail!("unknown daemon action"),
}
};
send_json(
&mut stream,
&Reply {
version: VERSION.into(),
workspace: identity.workspace.clone(),
state_root: identity.state_root.clone(),
status: status.into(),
},
)?;
if hello.action != "connect" || status != "ready" || stopping.load(Ordering::Acquire) {
return Ok(());
}
let connection = service.connect()?;
let result = relay(&mut stream, service, &connection, stopping);
let _ = service.disconnect(&connection);
result
}
fn relay(
stream: &mut UnixStream,
service: &PersistentService,
connection: &str,
stopping: &AtomicBool,
) -> Result<()> {
stream.set_nonblocking(true)?;
let mut input = Vec::new();
let mut frame_started = Instant::now();
let mut output = Vec::new();
let mut written = 0;
let mut output_started = Instant::now();
loop {
if stopping.load(Ordering::Acquire) {
return Ok(());
}
let mut buffer = [0; 8192];
match stream.read(&mut buffer) {
Ok(0) => return Ok(()),
Ok(count) => {
if input.is_empty() {
frame_started = Instant::now();
}
input.extend_from_slice(&buffer[..count]);
while let Some(end) = input.iter().position(|byte| *byte == b'\n') {
ensure!(end <= super::protocol::MAX_RECORD_BYTES, "frame too large");
service.submit(connection, &input[..end])?;
input.drain(..=end);
frame_started = Instant::now();
}
ensure!(
input.len() <= super::protocol::MAX_RECORD_BYTES,
"frame too large"
);
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {}
Err(error) => return Err(error.into()),
}
ensure!(
input.is_empty() || frame_started.elapsed() < TIMEOUT,
"incomplete frame timeout"
);
if output.is_empty()
&& let Some(record) = service.next_record(connection)?
{
output = record.into_bytes();
written = 0;
output_started = Instant::now();
}
if !output.is_empty() {
ensure!(
output_started.elapsed() < TIMEOUT,
"slow client output timeout"
);
match stream.write(&output[written..]) {
Ok(0) => bail!("client output closed"),
Ok(count) => written += count,
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {}
Err(error) => return Err(error.into()),
}
if written == output.len() {
let value: serde_json::Value = serde_json::from_slice(&output)?;
if value.get("kind").and_then(|kind| kind.as_str()) == Some("response")
&& let Some(id) = value.get("request_id").and_then(|id| id.as_str())
{
service.response_written(connection, id)?;
}
output.clear();
}
}
thread::sleep(Duration::from_millis(5));
}
}
#[cfg(test)]
mod survival_tests;