use std::ffi::OsString;
use std::io;
use std::os::unix::ffi::OsStringExt as _;
use std::os::unix::process::parent_id;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use tokio::time::{Instant, MissedTickBehavior, interval_at};
use crate::daemon;
use crate::environment::PromptEnvironment;
use crate::prompt::snapshot;
use crate::theme::AsyncTheme;
const REQUEST_MAGIC: &[u8] = b"ZTREQ";
const REQUEST_VERSION: &[u8] = b"1";
const PARENT_CHECK_INTERVAL: Duration = Duration::from_secs(1);
pub async fn serve_client(
instance: daemon::Instance,
shell_pid: u32,
theme: Arc<AsyncTheme>,
) -> io::Result<()> {
if parent_id() != shell_pid {
return Ok(());
}
let (sender, mut receiver) = mpsc::channel(4);
spawn_request_reader(sender);
let mut current: Option<JoinHandle<io::Result<()>>> = None;
let mut parent_check = interval_at(
Instant::now() + PARENT_CHECK_INTERVAL,
PARENT_CHECK_INTERVAL,
);
parent_check.set_missed_tick_behavior(MissedTickBehavior::Skip);
loop {
tokio::select! {
request = receiver.recv() => {
if let Some(request) = request {
if let Some(handle) = current.take() {
handle.abort();
if let Ok(Err(error)) = handle.await {
return Err(error);
}
}
let instance = instance.clone();
let theme = Arc::clone(&theme);
let environment = Arc::new(request.environment);
current = Some(tokio::spawn(async move {
snapshot(
request.generation,
request.cwd,
instance,
environment,
&theme,
)
.await
}));
} else {
cancel_current(&mut current).await;
break;
}
}
_ = parent_check.tick() => {
if parent_id() != shell_pid {
cancel_current(&mut current).await;
return Ok(());
}
}
}
}
Ok(())
}
async fn cancel_current(current: &mut Option<JoinHandle<io::Result<()>>>) {
if let Some(handle) = current.take() {
handle.abort();
let _ = handle.await;
}
}
fn spawn_request_reader(sender: mpsc::Sender<Request>) {
std::thread::Builder::new()
.name("ztheme-client-requests".into())
.spawn(move || {
let mut reader = std::io::BufReader::new(std::io::stdin());
loop {
match read_request(&mut reader) {
Ok(Some(request)) => {
if sender.blocking_send(request).is_err() {
return;
}
}
Ok(None) => return,
Err(error) => {
eprintln!("ztheme: client daemon request failed: {error}");
return;
}
}
}
})
.expect("spawning the request reader thread cannot fail");
}
struct Request {
generation: u64,
cwd: PathBuf,
environment: PromptEnvironment,
}
fn read_request<R>(reader: &mut R) -> io::Result<Option<Request>>
where
R: std::io::BufRead,
{
let magic = read_field(reader)?;
let Some(magic) = magic else {
return Ok(None);
};
if magic != REQUEST_MAGIC {
return Err(invalid_data("client request magic is invalid"));
}
let version = read_field(reader)?.ok_or_else(truncated)?;
if version != REQUEST_VERSION {
return Err(invalid_data("client request version is unsupported"));
}
let generation = read_field(reader)?.ok_or_else(truncated)?;
let generation = std::str::from_utf8(&generation)
.ok()
.and_then(|value| value.parse().ok())
.ok_or_else(|| invalid_data("client request generation is invalid"))?;
let cwd = read_field(reader)?.ok_or_else(truncated)?;
let cwd = PathBuf::from(OsString::from_vec(cwd));
if !cwd.is_absolute() {
return Err(invalid_data("client request cwd is not absolute"));
}
let environment = PromptEnvironment {
path: env_field(read_field(reader)?)?,
home: env_field(read_field(reader)?)?,
git_dir: env_field(read_field(reader)?)?,
git_work_tree: env_field(read_field(reader)?)?,
git_ceilings: env_field(read_field(reader)?)?,
virtual_env: env_field(read_field(reader)?)?,
conda_prefix: env_field(read_field(reader)?)?,
conda_default_env: env_field(read_field(reader)?)?,
perlbrew_perl: env_field(read_field(reader)?)?,
plenv_version: env_field(read_field(reader)?)?,
rustup_toolchain: env_field(read_field(reader)?)?,
rbenv_version: env_field(read_field(reader)?)?,
ruby_version: env_field(read_field(reader)?)?,
};
Ok(Some(Request {
generation,
cwd,
environment,
}))
}
fn read_field<R>(reader: &mut R) -> io::Result<Option<Vec<u8>>>
where
R: std::io::BufRead,
{
let mut field = Vec::with_capacity(64);
if reader.read_until(0, &mut field)? == 0 {
return Ok(None);
}
if field.pop() != Some(0) {
return Err(invalid_data("client request field is not NUL-terminated"));
}
Ok(Some(field))
}
fn env_field(field: Option<Vec<u8>>) -> io::Result<Option<OsString>> {
match field {
Some(value) if value.is_empty() => Ok(None),
Some(value) => Ok(Some(OsString::from_vec(value))),
None => Err(truncated()),
}
}
fn truncated() -> io::Error {
invalid_data("client request is truncated")
}
fn invalid_data(message: &'static str) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, message)
}
#[cfg(test)]
mod tests {
use std::ffi::OsStr;
use std::os::unix::ffi::OsStrExt as _;
use std::path::Path;
use std::io::BufReader;
use super::{REQUEST_MAGIC, REQUEST_VERSION, read_request};
const ENV_FIELD_COUNT: usize = 13;
fn request(cwd: &[u8], fields: &[&[u8]]) -> Vec<u8> {
assert_eq!(fields.len(), ENV_FIELD_COUNT);
let mut bytes = REQUEST_MAGIC.to_vec();
bytes.push(0);
bytes.extend_from_slice(REQUEST_VERSION);
bytes.push(0);
bytes.extend_from_slice(b"42");
bytes.push(0);
bytes.extend_from_slice(cwd);
bytes.push(0);
for field in fields {
bytes.extend_from_slice(field);
bytes.push(0);
}
bytes
}
#[test]
fn parses_a_complete_request_with_environment() {
let fields: [&[u8]; ENV_FIELD_COUNT] = [
b"/opt/bin:/usr/bin",
b"/home/user",
b"",
b"/work/tree",
b"",
b"",
b"",
b"",
b"",
b"",
b"",
b"",
b"",
];
let bytes = request(b"/work/project", &fields);
let mut reader = BufReader::new(&bytes[..]);
let request = read_request(&mut reader).unwrap().unwrap();
assert_eq!(request.generation, 42);
assert_eq!(request.cwd, Path::new("/work/project"));
assert_eq!(
request.environment.path.as_deref(),
Some(OsStr::new("/opt/bin:/usr/bin"))
);
assert_eq!(
request.environment.home.as_deref(),
Some(OsStr::new("/home/user"))
);
assert_eq!(request.environment.git_dir, None);
assert_eq!(
request.environment.git_work_tree.as_deref(),
Some(OsStr::new("/work/tree"))
);
assert!(read_request(&mut reader).unwrap().is_none());
}
#[test]
fn non_utf8_cwd_and_environment_round_trip() {
let mut git_dir = b"/repo-".to_vec();
git_dir.push(0xff);
let fields: [&[u8]; ENV_FIELD_COUNT] = [
b"", b"", &git_dir, b"", b"", b"", b"", b"", b"", b"", b"", b"", b"",
];
let mut cwd = b"/cwd-".to_vec();
cwd.push(0xfe);
let bytes = request(&cwd, &fields);
let mut reader = BufReader::new(&bytes[..]);
let request = read_request(&mut reader).unwrap().unwrap();
assert_eq!(request.cwd.as_os_str(), OsStr::from_bytes(&cwd));
assert_eq!(
request.environment.git_dir.as_deref(),
Some(OsStr::from_bytes(&git_dir))
);
}
#[test]
fn malformed_requests_are_rejected() {
let empty: [&[u8]; ENV_FIELD_COUNT] = [b""; ENV_FIELD_COUNT];
let valid = request(b"/work", &empty);
let mut bad_magic = valid.clone();
bad_magic[0] = b'X';
assert!(read_request(&mut BufReader::new(&bad_magic[..])).is_err());
let mut bad_version = valid.clone();
bad_version[6] = b'2';
assert!(read_request(&mut BufReader::new(&bad_version[..])).is_err());
let mut bad_generation = valid.clone();
bad_generation[8] = b'x';
assert!(read_request(&mut BufReader::new(&bad_generation[..])).is_err());
let mut relative_cwd = valid.clone();
relative_cwd[11] = b'.';
assert!(read_request(&mut BufReader::new(&relative_cwd[..])).is_err());
let truncated = &valid[..valid.len() - 3];
assert!(read_request(&mut BufReader::new(truncated)).is_err());
}
#[test]
fn clean_eof_is_a_normal_stop_but_partial_requests_are_rejected() {
assert!(
read_request(&mut BufReader::new(&b""[..]))
.unwrap()
.is_none()
);
let partial = b"ZTREQ\0";
assert!(read_request(&mut BufReader::new(&partial[..])).is_err());
}
}