use crate::daemon::sigpipe;
use crate::daemon::trust;
use anyhow::{anyhow, Context, Result};
use std::path::Path;
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::UnixStream;
const CONNECT_TIMEOUT: Duration = Duration::from_secs(1);
const CONTROL_TIMEOUT: Duration = Duration::from_secs(5);
const CONTROL_MAX_FRAME_BYTES: u32 = 256 * 1024;
const SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(45);
#[derive(Debug)]
pub(crate) enum ControlError {
Absent(anyhow::Error),
Untrusted(anyhow::Error),
Unintelligible(anyhow::Error),
}
impl ControlError {
pub(crate) fn kind(&self) -> &'static str {
match self {
Self::Absent(_) => "absent",
Self::Untrusted(_) => "untrusted",
Self::Unintelligible(_) => "unintelligible",
}
}
pub(crate) fn is_absent(&self) -> bool {
matches!(self, Self::Absent(_))
}
pub(crate) fn into_error(self) -> anyhow::Error {
match self {
Self::Absent(error) | Self::Untrusted(error) | Self::Unintelligible(error) => error,
}
}
}
impl std::fmt::Display for ControlError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Absent(error) | Self::Untrusted(error) | Self::Unintelligible(error) => {
write!(f, "{error:#}")
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct DaemonIdentity {
pub wire_version: Option<u64>,
pub keyhog_version: Option<String>,
}
impl std::fmt::Display for DaemonIdentity {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match (self.wire_version, self.keyhog_version.as_deref()) {
(Some(wire), Some(version)) => write!(f, "wire {wire}, keyhog {version}"),
(Some(wire), None) => write!(f, "wire {wire}, keyhog version undeclared"),
(None, Some(version)) => write!(f, "wire version undeclared, keyhog {version}"),
(None, None) => write!(f, "no declared wire or package version"),
}
}
}
pub(crate) struct ControlChannel {
stream: UnixStream,
_sigpipe: sigpipe::SigPipeGuard,
}
impl ControlChannel {
pub(crate) async fn connect(
socket_path: &Path,
) -> std::result::Result<(Self, DaemonIdentity), ControlError> {
trust::validate_socket_for_connect(socket_path).map_err(|error| {
if socket_path.exists() {
ControlError::Untrusted(error)
} else {
ControlError::Absent(error)
}
})?;
let stream = tokio::time::timeout(CONNECT_TIMEOUT, UnixStream::connect(socket_path))
.await
.map_err(|_| {
ControlError::Unintelligible(anyhow!(
"daemon control: connect to {} timed out after {}s",
socket_path.display(),
CONNECT_TIMEOUT.as_secs()
))
})?
.map_err(|error| {
let context = anyhow::Error::new(error).context(format!(
"daemon control: connect to {}",
socket_path.display()
));
ControlError::Absent(context)
})?;
trust::verify_connected_peer(&stream, socket_path).map_err(ControlError::Untrusted)?;
let mut channel = Self {
stream,
_sigpipe: sigpipe::SigPipeGuard::acquire(),
};
let hello = channel.round_trip("hello", CONTROL_TIMEOUT).await?;
let kind = hello.get("kind").and_then(serde_json::Value::as_str);
if kind != Some("hello") {
return Err(ControlError::Unintelligible(anyhow!(
"daemon control: expected a hello reply from {}, got kind {}",
socket_path.display(),
kind.unwrap_or("<absent>") )));
}
let identity = DaemonIdentity {
wire_version: hello
.get("wire_version")
.and_then(serde_json::Value::as_u64),
keyhog_version: hello
.get("keyhog_version")
.and_then(serde_json::Value::as_str)
.map(str::to_owned),
};
Ok((channel, identity))
}
pub(crate) async fn shutdown(&mut self) -> std::result::Result<(), ControlError> {
let response = self.round_trip("shutdown", SHUTDOWN_TIMEOUT).await?;
match response.get("kind").and_then(serde_json::Value::as_str) {
Some("shutdown") => Ok(()),
other => Err(ControlError::Unintelligible(anyhow!(
"daemon control: shutdown was not acknowledged (reply kind {})",
other.unwrap_or("<absent>") ))),
}
}
async fn round_trip(
&mut self,
op: &str,
timeout: Duration,
) -> std::result::Result<serde_json::Value, ControlError> {
let request = serde_json::json!({ "op": op });
tokio::time::timeout(timeout, async {
write_frame(&mut self.stream, &request).await?;
read_frame(&mut self.stream).await
})
.await
.map_err(|_| {
ControlError::Unintelligible(anyhow!(
"daemon control: {op} did not complete within {}s",
timeout.as_secs()
))
})?
.map_err(ControlError::Unintelligible)
}
}
async fn write_frame(stream: &mut UnixStream, value: &serde_json::Value) -> Result<()> {
let body = serde_json::to_vec(value).context("daemon control: encode request")?;
let length = u32::try_from(body.len())
.map_err(|_| anyhow!("daemon control: request body exceeds the frame length prefix"))?;
stream
.write_all(&length.to_be_bytes())
.await
.context("daemon control: write frame length")?;
stream
.write_all(&body)
.await
.context("daemon control: write frame body")?;
stream
.flush()
.await
.context("daemon control: flush frame")?;
Ok(())
}
async fn read_frame(stream: &mut UnixStream) -> Result<serde_json::Value> {
let mut length = [0u8; 4];
stream
.read_exact(&mut length)
.await
.context("daemon control: read frame length")?;
let length = u32::from_be_bytes(length);
if length > CONTROL_MAX_FRAME_BYTES {
anyhow::bail!(
"daemon control: peer announced a {length} byte administration frame, above the \
{CONTROL_MAX_FRAME_BYTES} byte ceiling"
);
}
let mut body = vec![0u8; length as usize];
stream
.read_exact(&mut body)
.await
.context("daemon control: read frame body")?;
serde_json::from_slice(&body).context("daemon control: parse reply")
}