1use crate::protocol::{DaemonInfo, Request, Response, PROTOCOL_VERSION};
11use crate::warm::WarmState;
12use pushkin_core::manifest::Manifest;
13use pushkin_core::pipeline::WriteRequest;
14use std::path::{Path, PathBuf};
15use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
16use tokio::net::{UnixListener, UnixStream};
17use tokio::task::JoinSet;
18
19pub const SOCKET_FILE: &str = ".pushkin/daemon.sock";
22
23const MAX_LINE_BYTES: usize = 16 * 1024 * 1024;
27
28#[derive(Debug, thiserror::Error)]
29pub enum ServerError {
30 #[error("daemon io failure: {0}")]
31 Io(#[from] std::io::Error),
32 #[error("daemon not running")]
33 NotRunning,
34 #[error("protocol failure: {0}")]
35 Protocol(String),
36}
37
38#[must_use]
39pub fn socket_path(repo_root: &Path) -> PathBuf {
40 repo_root.join(SOCKET_FILE)
41}
42
43pub fn serve(repo_root: &Path, manifest: Manifest) -> Result<(), ServerError> {
52 serve_at(&socket_path(repo_root), manifest, false)
53}
54
55pub fn serve_at(socket: &Path, manifest: Manifest, read_only: bool) -> Result<(), ServerError> {
64 serve_resolved(socket, manifest, read_only, None)
65}
66
67pub fn serve_resolved(
83 socket: &Path,
84 manifest: Manifest,
85 read_only: bool,
86 governing: Option<&Path>,
87) -> Result<(), ServerError> {
88 let warm = WarmState::new(manifest);
89 let watch_root = socket
94 .parent()
95 .and_then(Path::parent)
96 .filter(|root| !root.as_os_str().is_empty())
97 .unwrap_or_else(|| Path::new("."));
98 let governing = governing.map_or_else(|| watch_root.join("pushkin.toml"), Path::to_path_buf);
99 let _watch_guard = warm.watch_governing(watch_root, &governing).ok();
100 if let Some(parent) = socket.parent() {
101 std::fs::create_dir_all(parent)?;
102 }
103 if socket.exists() {
105 std::fs::remove_file(socket)?;
106 }
107
108 let runtime = tokio::runtime::Builder::new_current_thread()
109 .enable_io()
110 .enable_time()
111 .build()
112 .map_err(|e| ServerError::Protocol(format!("runtime build failed: {e}")))?;
113
114 let result = runtime.block_on(serve_inner(socket, &warm, read_only));
115 let _ = std::fs::remove_file(socket);
117 result
118}
119
120async fn serve_inner(socket: &Path, warm: &WarmState, read_only: bool) -> Result<(), ServerError> {
121 let listener = UnixListener::bind(socket)?;
122 let (shutdown_tx, mut shutdown_rx) = tokio::sync::mpsc::channel::<()>(1);
123 let mut connections: JoinSet<()> = JoinSet::new();
124
125 loop {
126 tokio::select! {
127 accepted = listener.accept() => {
128 let Ok((stream, _addr)) = accepted else { continue };
129 let warm = warm.share();
130 let shutdown_tx = shutdown_tx.clone();
131 connections.spawn(async move {
132 let _ = handle_connection(stream, &warm, read_only, &shutdown_tx).await;
135 });
136 }
137 _ = shutdown_rx.recv() => break,
138 Some(_) = connections.join_next(), if !connections.is_empty() => {}
140 }
141 }
142 while connections.join_next().await.is_some() {}
144 Ok(())
145}
146
147async fn handle_connection(
148 stream: UnixStream,
149 warm: &WarmState,
150 read_only: bool,
151 shutdown_tx: &tokio::sync::mpsc::Sender<()>,
152) -> Result<(), ServerError> {
153 let (read_half, mut write_half) = stream.into_split();
154 let mut lines = BufReader::with_capacity(64 * 1024, read_half).lines();
155
156 while let Ok(Some(line)) = lines.next_line().await {
157 if line.len() > MAX_LINE_BYTES {
158 let response = Response::Error {
159 message: "request exceeds size cap".to_owned(),
160 };
161 write_response(&mut write_half, &response).await?;
162 continue;
163 }
164 let response = match serde_json::from_str::<Request>(&line) {
165 Ok(request) => {
166 let response = respond(&request, warm, read_only);
167 let stop_serving = !read_only && matches!(request, Request::Shutdown { .. });
168 write_response(&mut write_half, &response).await?;
169 if stop_serving {
170 let _ = shutdown_tx.send(()).await;
171 return Ok(());
172 }
173 continue;
174 }
175 Err(error) => Response::Error {
176 message: format!("unrecognized request: {error}"),
177 },
178 };
179 write_response(&mut write_half, &response).await?;
180 }
181 Ok(())
182}
183
184async fn write_response(
185 write_half: &mut tokio::net::unix::OwnedWriteHalf,
186 response: &Response,
187) -> Result<(), ServerError> {
188 let mut payload = serde_json::to_string(response)
189 .map_err(|e| ServerError::Protocol(format!("response encode failed: {e}")))?;
190 payload.push('\n');
191 write_half.write_all(payload.as_bytes()).await?;
192 write_half.flush().await?;
193 Ok(())
194}
195
196fn respond(request: &Request, warm: &WarmState, read_only: bool) -> Response {
199 let v = match request {
200 Request::Check { v, .. } | Request::Ping { v } | Request::Shutdown { v } => *v,
201 };
202 if v != PROTOCOL_VERSION {
203 return Response::Error {
204 message: format!("protocol version {v} unsupported (daemon speaks {PROTOCOL_VERSION})"),
205 };
206 }
207 match request {
208 Request::Check {
209 file_path, content, ..
210 } => Response::Check {
211 result: warm.check(&WriteRequest {
212 file_path: file_path.clone(),
213 content: content.clone(),
214 }),
215 },
216 Request::Ping { .. } => Response::Pong {
217 info: DaemonInfo {
218 pid: std::process::id(),
219 version: env!("CARGO_PKG_VERSION").to_owned(),
220 read_only,
221 },
222 },
223 Request::Shutdown { .. } => {
224 if read_only {
225 Response::Error {
226 message: "read-only daemon: lifecycle mutations are refused over the wire \
227 (kill the process from the session that spawned it)"
228 .to_owned(),
229 }
230 } else {
231 Response::ShuttingDown
232 }
233 }
234 }
235}
236
237pub fn request(repo_root: &Path, request: &Request) -> Result<Response, ServerError> {
244 request_at(&socket_path(repo_root), request)
245}
246
247pub fn request_at(socket: &Path, request: &Request) -> Result<Response, ServerError> {
255 use std::io::{BufRead, BufReader as StdBufReader, Write};
256
257 let mut stream = match std::os::unix::net::UnixStream::connect(socket) {
258 Ok(stream) => stream,
259 Err(error) => {
260 return Err(match error.kind() {
261 std::io::ErrorKind::NotFound | std::io::ErrorKind::ConnectionRefused => {
262 ServerError::NotRunning
263 }
264 _ => ServerError::Io(error),
265 })
266 }
267 };
268 let mut payload = serde_json::to_string(request)
269 .map_err(|e| ServerError::Protocol(format!("request encode failed: {e}")))?;
270 payload.push('\n');
271 stream.write_all(payload.as_bytes())?;
272 stream.flush()?;
273
274 let mut line = String::new();
275 StdBufReader::new(&mut stream).read_line(&mut line)?;
276 if line.is_empty() {
277 return Err(ServerError::Protocol(
278 "daemon closed the connection without responding".to_owned(),
279 ));
280 }
281 serde_json::from_str(&line)
282 .map_err(|e| ServerError::Protocol(format!("unrecognized response: {e}")))
283}