use crate::Result;
use crate::daemon::{Daemon, RunOptions};
use crate::daemon_id::DaemonId;
use crate::env;
#[cfg(unix)]
use crate::error::IpcError;
use interprocess::local_socket::Name;
#[cfg(unix)]
use interprocess::local_socket::{GenericFilePath, ToFsName};
#[cfg(windows)]
use interprocess::local_socket::{GenericNamespaced, ToNsName};
use miette::{Context, IntoDiagnostic};
#[cfg(unix)]
use std::path::Path;
use std::path::PathBuf;
pub(crate) mod batch;
pub(crate) mod client;
pub(crate) mod server;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, strum::Display, strum::EnumIs)]
#[allow(clippy::large_enum_variant)]
pub enum IpcRequest {
Connect,
ConnectV2 {
version: String,
},
Clean,
Stop {
id: DaemonId,
},
GetActiveDaemons,
GetDisabledDaemons,
Run(RunOptions),
Enable {
id: DaemonId,
},
Disable {
id: DaemonId,
},
UpdateShellDir {
shell_pid: u32,
dir: PathBuf,
},
GetNotifications,
SyncMdns,
ReloadConfig,
ProjectEnter {
pid: u32,
dir: PathBuf,
},
ProjectLeave {
pid: u32,
dir: PathBuf,
},
GetProjectSessions,
SinkOutputLine {
id: DaemonId,
token: u64,
fires_hook: bool,
line: String,
},
GetWebUrl,
CleanFiltered {
namespaces: Vec<String>,
daemons: Vec<DaemonId>,
prune: bool,
},
ClaimDaemons {
ids: Vec<DaemonId>,
},
#[serde(skip)]
Invalid {
error: String,
},
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ProjectSessionInfo {
pub pid: u32,
pub directory: PathBuf,
#[serde(skip_serializing_if = "Option::is_none", default)]
pub liveness_title: Option<String>,
pub alive: bool,
#[serde(skip_serializing_if = "Option::is_none", default)]
pub current_title: Option<String>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, strum::Display, strum::EnumIs)]
pub enum IpcResponse {
Ok,
ConnectOk {
version: String,
},
Yes,
No,
Error(String),
Notifications(Vec<(log::LevelFilter, String)>),
ActiveDaemons(Vec<Daemon>),
DisabledDaemons(Vec<DaemonId>),
DaemonAlreadyRunning,
DaemonStart {
daemon: Daemon,
},
DaemonFailed {
error: String,
},
PortConflict {
port: u16,
process: String,
pid: u32,
},
NoAvailablePort {
start_port: u16,
attempts: u32,
},
DaemonReady {
daemon: Daemon,
},
DaemonFailedWithCode {
exit_code: Option<i32>,
#[serde(default)]
resolved_ports: Vec<u16>,
},
DaemonWasNotRunning,
MdnsSynced,
ConfigReloaded,
WebUrl {
url: Option<String>,
},
DaemonStopFailed {
error: String,
},
DaemonNotRunning,
DaemonNotFound,
ProjectSessions(Vec<ProjectSessionInfo>),
Cleaned {
count: u64,
},
}
#[cfg(unix)]
const SOCKET_PATH_CAPACITY: usize = {
let sun = unsafe { std::mem::zeroed::<libc::sockaddr_un>() };
sun.sun_path.len()
};
#[cfg(unix)]
fn check_socket_path(path: &Path, capacity: usize) -> Result<()> {
use std::os::unix::ffi::OsStrExt;
let len = path.as_os_str().as_bytes().len();
if len <= capacity {
return Ok(());
}
let over = len - capacity;
let help = format!(
"the socket is {path} and is {over} byte(s) over the limit.\n\
Set PITCHFORK_STATE_DIR (or XDG_STATE_HOME) to a shorter directory, for example \
PITCHFORK_STATE_DIR=/tmp/pitchfork, so that <state dir>/sock/main.sock is at most \
{capacity} bytes.",
path = path.display()
);
Err(IpcError::SocketPathTooLong {
path: path.to_path_buf(),
len,
limit: capacity,
help,
}
.into())
}
fn fs_name(name: &str) -> Result<Name<'_>> {
#[cfg(unix)]
{
let path = env::IPC_SOCK_DIR.join(name).with_extension("sock");
check_socket_path(&path, SOCKET_PATH_CAPACITY)?;
let fs_name = path.to_fs_name::<GenericFilePath>().into_diagnostic()?;
Ok(fs_name)
}
#[cfg(windows)]
{
let state_dir = env::PITCHFORK_STATE_DIR.to_string_lossy();
let mut hash: u64 = 0xcbf29ce484222325;
for byte in state_dir.bytes() {
hash ^= byte as u64;
hash = hash.wrapping_mul(0x100000001b3);
}
let pipe_name = format!("pitchfork-{hash:016x}-{name}");
pipe_name
.to_ns_name::<GenericNamespaced>()
.into_diagnostic()
}
}
pub(crate) fn socket_display() -> String {
#[cfg(unix)]
{
env::IPC_SOCK_MAIN.display().to_string()
}
#[cfg(windows)]
{
"the supervisor named pipe".to_string()
}
}
pub(crate) async fn supervisor_listening() -> bool {
use interprocess::local_socket::traits::tokio::Stream as _;
let Ok(name) = fs_name("main") else {
return false;
};
let connect = interprocess::local_socket::tokio::Stream::connect(name);
match tokio::time::timeout(std::time::Duration::from_secs(1), connect).await {
Ok(Ok(_)) => true,
Ok(Err(err)) => {
trace!("no supervisor listening on the IPC socket: {err}");
false
}
Err(_) => {
debug!("timed out probing the IPC socket; treating it as not listening");
false
}
}
}
fn serialize<T: serde::Serialize>(msg: &T) -> Result<Vec<u8>> {
serde_json::to_vec(msg)
.into_diagnostic()
.wrap_err("failed to serialize IPC message as JSON")
}
fn deserialize<T: serde::de::DeserializeOwned>(bytes: &[u8]) -> Result<T> {
let mut bytes = bytes.to_vec();
bytes.pop();
let preview = std::str::from_utf8(&bytes).unwrap_or("<binary>");
trace!("msg: {preview:?}");
serde_json::from_slice(&bytes)
.into_diagnostic()
.wrap_err("failed to deserialize IPC JSON response")
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(unix)]
#[test]
fn socket_path_over_sun_path_capacity_names_path_length_and_fix() {
let fits = PathBuf::from(format!("/{}", "a".repeat(SOCKET_PATH_CAPACITY - 1)));
assert_eq!(fits.as_os_str().len(), SOCKET_PATH_CAPACITY);
assert!(check_socket_path(&fits, SOCKET_PATH_CAPACITY).is_ok());
let long = PathBuf::from(format!(
"/{}/sock/main.sock",
"a".repeat(SOCKET_PATH_CAPACITY)
));
let err = check_socket_path(&long, SOCKET_PATH_CAPACITY).unwrap_err();
let len = long.as_os_str().len();
let message = err.to_string();
assert!(message.contains(&format!("{len} bytes")), "{message}");
assert!(
message.contains(&format!("allows {SOCKET_PATH_CAPACITY}")),
"{message}"
);
let help = err.help().expect("help text").to_string();
assert!(help.contains(&long.display().to_string()), "{help}");
assert!(help.contains("PITCHFORK_STATE_DIR"), "{help}");
assert!(help.contains("XDG_STATE_HOME"), "{help}");
assert!(
matches!(
err.downcast_ref::<IpcError>(),
Some(IpcError::SocketPathTooLong { len: l, limit, .. })
if *l == len && *limit == SOCKET_PATH_CAPACITY
),
"{err:?}"
);
}
#[cfg(unix)]
#[test]
fn socket_path_capacity_matches_the_platform() {
assert!((100..=108).contains(&SOCKET_PATH_CAPACITY));
}
#[test]
fn filtered_clean_ipc_round_trips() {
let request = IpcRequest::CleanFiltered {
namespaces: vec!["worktree".to_string()],
daemons: vec![DaemonId::new("worktree", "api")],
prune: true,
};
let mut bytes = serialize(&request).unwrap();
bytes.push(b'\n');
let decoded: IpcRequest = deserialize(&bytes).unwrap();
match decoded {
IpcRequest::CleanFiltered {
namespaces,
daemons,
prune,
} => {
assert_eq!(namespaces, ["worktree"]);
assert_eq!(daemons, [DaemonId::new("worktree", "api")]);
assert!(prune);
}
other => panic!("unexpected request: {other:?}"),
}
let mut bytes = serialize(&IpcResponse::Cleaned { count: 3 }).unwrap();
bytes.push(b'\n');
let decoded: IpcResponse = deserialize(&bytes).unwrap();
assert!(matches!(decoded, IpcResponse::Cleaned { count: 3 }));
}
fn round_trip<T: serde::Serialize + serde::de::DeserializeOwned>(value: &T) -> T {
let mut bytes = serialize(value).unwrap();
bytes.push(b'\n');
deserialize(&bytes).unwrap()
}
#[test]
fn claim_daemons_ipc_round_trips() {
let ids = vec![DaemonId::new("proj", "api"), DaemonId::new("proj", "db")];
match round_trip(&IpcRequest::ClaimDaemons { ids: ids.clone() }) {
IpcRequest::ClaimDaemons { ids: decoded } => assert_eq!(decoded, ids),
other => panic!("unexpected request: {other:?}"),
}
}
#[test]
fn proxy_idle_timeout_survives_the_ipc_encoding() {
let daemon = Daemon {
id: DaemonId::new("proj", "api"),
proxy_idle_timeout_ms: Some(900_000),
..Default::default()
};
match round_trip(&IpcResponse::ActiveDaemons(vec![daemon])) {
IpcResponse::ActiveDaemons(daemons) => {
assert_eq!(daemons[0].proxy_idle_timeout_ms, Some(900_000));
assert!(!daemons[0].oneshot);
}
other => panic!("unexpected response: {other:?}"),
}
let opts = RunOptions {
id: DaemonId::new("proj", "api"),
proxy_idle_timeout_ms: Some(900_000),
..Default::default()
};
match round_trip(&IpcRequest::Run(opts)) {
IpcRequest::Run(opts) => assert_eq!(opts.proxy_idle_timeout_ms, Some(900_000)),
other => panic!("unexpected request: {other:?}"),
}
}
}