use std::fs::File;
use std::path::Path;
use std::process::{Command, Stdio};
use std::time::{Duration, Instant};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufReader};
use tokio::net::UnixStream;
use tokio::net::unix::{OwnedReadHalf, OwnedWriteHalf};
use super::handshake::{
ControlCommand, Hello, HelloMode, HelloReply, read_json_line, write_json_line,
};
use super::paths::DaemonPaths;
use super::prepare_run_dir;
use super::server::daemon_running;
const REPLY_TIMEOUT: Duration = Duration::from_secs(30);
const MAX_LOG_BYTES: u64 = 5 * 1024 * 1024;
const LOG_TAIL_LINES: usize = 20;
pub struct DaemonConnection {
pub reader: BufReader<OwnedReadHalf>,
pub writer: OwnedWriteHalf,
pub reply: HelloReply,
}
pub async fn connect(paths: &DaemonPaths) -> std::io::Result<UnixStream> {
if paths.socket_is_external() {
verify_socket_owner(paths)?;
}
UnixStream::connect(&paths.socket).await
}
fn verify_socket_owner(paths: &DaemonPaths) -> std::io::Result<()> {
use std::os::unix::fs::MetadataExt;
let socket_uid = std::fs::symlink_metadata(&paths.socket)?.uid();
let own_uid = std::fs::metadata(paths.run_dir())?.uid();
if socket_uid != own_uid {
return Err(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
format!(
"{} is owned by another user; refusing to connect",
paths.socket.display()
),
));
}
Ok(())
}
pub async fn connect_or_spawn<F>(
paths: &DaemonPaths,
spawn: F,
startup_timeout: Duration,
) -> anyhow::Result<UnixStream>
where
F: FnOnce(&DaemonPaths) -> anyhow::Result<()>,
{
if let Ok(stream) = connect(paths).await {
return Ok(stream);
}
prepare_run_dir(paths.run_dir())?;
let lock_path = paths.spawn_lock.clone();
let _spawn_lock = tokio::task::spawn_blocking(move || -> std::io::Result<File> {
let file = File::options()
.write(true)
.create(true)
.truncate(false)
.open(lock_path)?;
file.lock()?;
Ok(file)
})
.await??;
if let Ok(stream) = connect(paths).await {
return Ok(stream);
}
if !daemon_running(paths) {
let _ = std::fs::remove_file(&paths.socket);
spawn(paths)?;
}
let deadline = Instant::now() + startup_timeout;
loop {
match connect(paths).await {
Ok(stream) => return Ok(stream),
Err(_) if Instant::now() < deadline => {
tokio::time::sleep(Duration::from_millis(50)).await;
}
Err(e) => anyhow::bail!(
"Scryer daemon did not accept connections on {} within {startup_timeout:?} ({e}).{}",
paths.socket.display(),
log_tail(&paths.log_file)
),
}
}
}
const STALE_STOP_TIMEOUT: Duration = Duration::from_secs(10);
pub fn is_stale_build(build: Option<&str>) -> bool {
build != Some(crate::build_info::build_id())
}
pub async fn replace_stale_daemon(paths: &DaemonPaths) -> anyhow::Result<bool> {
let Some(reply) = control(paths, ControlCommand::Status).await? else {
return Ok(false);
};
let Some(status) = reply.status else {
return Ok(false);
};
if !is_stale_build(status.build.as_deref()) || status.sessions > 0 {
return Ok(false);
}
tracing::info!(
"Stopping Scryer daemon pid {} (build {}): this client is build {}",
status.pid,
status.build.as_deref().unwrap_or("unknown"),
crate::build_info::build_id()
);
let stopped = match control(paths, ControlCommand::StopIfIdle).await {
Ok(Some(reply)) => reply.status.is_none(),
Ok(None) => false,
Err(_) => matches!(control(paths, ControlCommand::Stop).await, Ok(Some(_))),
};
if !stopped {
tracing::info!(
"Stale Scryer daemon pid {} gained a session; leaving it",
status.pid
);
return Ok(false);
}
let deadline = Instant::now() + STALE_STOP_TIMEOUT;
while daemon_running(paths) && pid_file_owner(paths) == Some(status.pid) {
if Instant::now() >= deadline {
anyhow::bail!(
"stale Scryer daemon (pid {}) did not exit within {STALE_STOP_TIMEOUT:?}",
status.pid
);
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
Ok(true)
}
fn pid_file_owner(paths: &DaemonPaths) -> Option<u32> {
std::fs::read_to_string(&paths.pid_file)
.ok()?
.trim()
.parse()
.ok()
}
pub async fn handshake(stream: UnixStream, hello: &Hello) -> anyhow::Result<DaemonConnection> {
let (read, mut writer) = stream.into_split();
write_json_line(&mut writer, hello).await?;
let mut reader = BufReader::new(read);
let reply: HelloReply = tokio::time::timeout(REPLY_TIMEOUT, read_json_line(&mut reader))
.await
.map_err(|_| anyhow::anyhow!("timed out waiting for the Scryer daemon handshake"))??;
if !reply.ok {
anyhow::bail!(
"Scryer daemon rejected the connection: {}",
reply.error.as_deref().unwrap_or("unknown error")
);
}
Ok(DaemonConnection {
reader,
writer,
reply,
})
}
pub async fn control(
paths: &DaemonPaths,
cmd: ControlCommand,
) -> anyhow::Result<Option<HelloReply>> {
let Ok(stream) = connect(paths).await else {
return Ok(None);
};
let conn = handshake(stream, &Hello::new(HelloMode::Control { cmd })).await?;
Ok(Some(conn.reply))
}
pub fn spawn_daemon_process(
exe: &Path,
paths: &DaemonPaths,
idle_timeout_secs: u64,
) -> anyhow::Result<()> {
use std::os::unix::process::CommandExt;
prepare_run_dir(paths.run_dir())?;
if std::fs::metadata(&paths.log_file).is_ok_and(|m| m.len() > MAX_LOG_BYTES) {
let _ = std::fs::remove_file(&paths.log_file);
}
let log = File::options()
.create(true)
.append(true)
.open(&paths.log_file)?;
let mut child = Command::new(exe)
.arg("daemon")
.arg("run")
.arg("--db-url")
.arg(&paths.db_path)
.arg("--idle-timeout")
.arg(idle_timeout_secs.to_string())
.env_remove("SCRYER_DB_URL")
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(log)
.process_group(0)
.spawn()?;
std::thread::spawn(move || {
let _ = child.wait();
});
Ok(())
}
pub async fn pipe<I, O, R, W>(
mut input: I,
mut output: O,
mut daemon_read: R,
mut daemon_write: W,
) -> std::io::Result<()>
where
I: AsyncRead + Unpin,
O: AsyncWrite + Unpin,
R: AsyncRead + Unpin,
W: AsyncWrite + Unpin,
{
let upstream = async {
let copied = tokio::io::copy(&mut input, &mut daemon_write).await;
let _ = daemon_write.shutdown().await;
copied.map(|_| ())
};
let downstream = async {
let mut buf = vec![0u8; 64 * 1024];
loop {
let n = daemon_read.read(&mut buf).await?;
if n == 0 {
return Ok::<_, std::io::Error>(());
}
output.write_all(&buf[..n]).await?;
output.flush().await?;
}
};
tokio::pin!(upstream, downstream);
tokio::select! {
res = &mut downstream => res,
res = &mut upstream => {
res?;
downstream.await
}
}
}
fn log_tail(log_file: &Path) -> String {
let Ok(content) = std::fs::read_to_string(log_file) else {
return String::new();
};
let lines: Vec<&str> = content.lines().collect();
if lines.is_empty() {
return String::new();
}
let tail = &lines[lines.len().saturating_sub(LOG_TAIL_LINES)..];
format!(
"\nLast lines of {}:\n{}",
log_file.display(),
tail.join("\n")
)
}