use crate::protocol::{DaemonInfo, Request, Response, PROTOCOL_VERSION};
use crate::warm::WarmState;
use pushkin_core::manifest::Manifest;
use pushkin_core::pipeline::WriteRequest;
use std::path::{Path, PathBuf};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::{UnixListener, UnixStream};
use tokio::task::JoinSet;
pub const SOCKET_FILE: &str = ".pushkin/daemon.sock";
const MAX_LINE_BYTES: usize = 16 * 1024 * 1024;
#[derive(Debug, thiserror::Error)]
pub enum ServerError {
#[error("daemon io failure: {0}")]
Io(#[from] std::io::Error),
#[error("daemon not running")]
NotRunning,
#[error("protocol failure: {0}")]
Protocol(String),
}
#[must_use]
pub fn socket_path(repo_root: &Path) -> PathBuf {
repo_root.join(SOCKET_FILE)
}
pub fn serve(repo_root: &Path, manifest: Manifest) -> Result<(), ServerError> {
serve_at(&socket_path(repo_root), manifest, false)
}
pub fn serve_at(socket: &Path, manifest: Manifest, read_only: bool) -> Result<(), ServerError> {
serve_resolved(socket, manifest, read_only, None)
}
pub fn serve_resolved(
socket: &Path,
manifest: Manifest,
read_only: bool,
governing: Option<&Path>,
) -> Result<(), ServerError> {
let warm = WarmState::new(manifest);
let watch_root = socket
.parent()
.and_then(Path::parent)
.filter(|root| !root.as_os_str().is_empty())
.unwrap_or_else(|| Path::new("."));
let governing = governing.map_or_else(|| watch_root.join("pushkin.toml"), Path::to_path_buf);
let _watch_guard = warm.watch_governing(watch_root, &governing).ok();
if let Some(parent) = socket.parent() {
std::fs::create_dir_all(parent)?;
}
if socket.exists() {
std::fs::remove_file(socket)?;
}
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_io()
.enable_time()
.build()
.map_err(|e| ServerError::Protocol(format!("runtime build failed: {e}")))?;
let result = runtime.block_on(serve_inner(socket, &warm, read_only));
let _ = std::fs::remove_file(socket);
result
}
async fn serve_inner(socket: &Path, warm: &WarmState, read_only: bool) -> Result<(), ServerError> {
let listener = UnixListener::bind(socket)?;
let (shutdown_tx, mut shutdown_rx) = tokio::sync::mpsc::channel::<()>(1);
let mut connections: JoinSet<()> = JoinSet::new();
loop {
tokio::select! {
accepted = listener.accept() => {
let Ok((stream, _addr)) = accepted else { continue };
let warm = warm.share();
let shutdown_tx = shutdown_tx.clone();
connections.spawn(async move {
let _ = handle_connection(stream, &warm, read_only, &shutdown_tx).await;
});
}
_ = shutdown_rx.recv() => break,
Some(_) = connections.join_next(), if !connections.is_empty() => {}
}
}
while connections.join_next().await.is_some() {}
Ok(())
}
async fn handle_connection(
stream: UnixStream,
warm: &WarmState,
read_only: bool,
shutdown_tx: &tokio::sync::mpsc::Sender<()>,
) -> Result<(), ServerError> {
let (read_half, mut write_half) = stream.into_split();
let mut lines = BufReader::with_capacity(64 * 1024, read_half).lines();
while let Ok(Some(line)) = lines.next_line().await {
if line.len() > MAX_LINE_BYTES {
let response = Response::Error {
message: "request exceeds size cap".to_owned(),
};
write_response(&mut write_half, &response).await?;
continue;
}
let response = match serde_json::from_str::<Request>(&line) {
Ok(request) => {
let response = respond(&request, warm, read_only);
let stop_serving = !read_only && matches!(request, Request::Shutdown { .. });
write_response(&mut write_half, &response).await?;
if stop_serving {
let _ = shutdown_tx.send(()).await;
return Ok(());
}
continue;
}
Err(error) => Response::Error {
message: format!("unrecognized request: {error}"),
},
};
write_response(&mut write_half, &response).await?;
}
Ok(())
}
async fn write_response(
write_half: &mut tokio::net::unix::OwnedWriteHalf,
response: &Response,
) -> Result<(), ServerError> {
let mut payload = serde_json::to_string(response)
.map_err(|e| ServerError::Protocol(format!("response encode failed: {e}")))?;
payload.push('\n');
write_half.write_all(payload.as_bytes()).await?;
write_half.flush().await?;
Ok(())
}
fn respond(request: &Request, warm: &WarmState, read_only: bool) -> Response {
let v = match request {
Request::Check { v, .. } | Request::Ping { v } | Request::Shutdown { v } => *v,
};
if v != PROTOCOL_VERSION {
return Response::Error {
message: format!("protocol version {v} unsupported (daemon speaks {PROTOCOL_VERSION})"),
};
}
match request {
Request::Check {
file_path, content, ..
} => Response::Check {
result: warm.check(&WriteRequest {
file_path: file_path.clone(),
content: content.clone(),
}),
},
Request::Ping { .. } => Response::Pong {
info: DaemonInfo {
pid: std::process::id(),
version: env!("CARGO_PKG_VERSION").to_owned(),
read_only,
},
},
Request::Shutdown { .. } => {
if read_only {
Response::Error {
message: "read-only daemon: lifecycle mutations are refused over the wire \
(kill the process from the session that spawned it)"
.to_owned(),
}
} else {
Response::ShuttingDown
}
}
}
}
pub fn request(repo_root: &Path, request: &Request) -> Result<Response, ServerError> {
request_at(&socket_path(repo_root), request)
}
pub fn request_at(socket: &Path, request: &Request) -> Result<Response, ServerError> {
use std::io::{BufRead, BufReader as StdBufReader, Write};
let mut stream = match std::os::unix::net::UnixStream::connect(socket) {
Ok(stream) => stream,
Err(error) => {
return Err(match error.kind() {
std::io::ErrorKind::NotFound | std::io::ErrorKind::ConnectionRefused => {
ServerError::NotRunning
}
_ => ServerError::Io(error),
})
}
};
let mut payload = serde_json::to_string(request)
.map_err(|e| ServerError::Protocol(format!("request encode failed: {e}")))?;
payload.push('\n');
stream.write_all(payload.as_bytes())?;
stream.flush()?;
let mut line = String::new();
StdBufReader::new(&mut stream).read_line(&mut line)?;
if line.is_empty() {
return Err(ServerError::Protocol(
"daemon closed the connection without responding".to_owned(),
));
}
serde_json::from_str(&line)
.map_err(|e| ServerError::Protocol(format!("unrecognized response: {e}")))
}