pub const PROTOCOL_VERSION: u32 = 1;
#[derive(Debug, Clone)]
pub struct StatusInfo {
pub pid: u32,
pub socket: String,
pub vaults: Vec<String>,
pub idle_secs: u64,
pub ttl_secs: u64,
}
pub fn ttl_secs() -> u64 {
const DEFAULT: u64 = 300;
match std::env::var_os("ZKV_LOCK_SECS").map(|raw| raw.to_string_lossy().parse::<u64>()) {
Some(Ok(v)) => v,
_ => DEFAULT,
}
}
pub fn enabled() -> bool {
cfg!(unix)
&& ttl_secs() > 0
&& std::env::var_os("ZKV_NO_AGENT").is_none_or(|v| v != "1")
}
#[cfg(unix)]
mod imp {
use super::{StatusInfo, PROTOCOL_VERSION};
use crate::crypto::{KdfParams, MasterKey};
use crate::error::Result;
use serde::{de::DeserializeOwned, Deserialize, Serialize};
use std::collections::HashMap;
use std::fs;
use std::io::{BufRead, BufReader, BufWriter, Write};
use std::os::unix::fs::DirBuilderExt;
use std::os::unix::fs::PermissionsExt;
use std::os::unix::net::{UnixListener, UnixStream};
use std::os::unix::process::CommandExt;
use std::path::{Path, PathBuf};
use std::process::{Command, Stdio};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::{Duration, Instant};
use zeroize::Zeroizing;
#[derive(Serialize, Deserialize, Debug)]
struct Handshake {
proto: u32,
}
#[derive(Serialize, Deserialize, Debug)]
enum Request {
Get { path: String },
Put {
path: String,
key: [u8; 32],
kdf: KdfParams,
salt: [u8; 16],
},
Forget { path: String },
Status,
Stop,
Lock,
}
#[derive(Serialize, Deserialize, Debug)]
enum Response {
Got {
key: [u8; 32],
kdf: KdfParams,
salt: [u8; 16],
},
Miss,
Ok,
StatusResp {
pid: u32,
socket: String,
vaults: Vec<String>,
idle_secs: u64,
ttl_secs: u64,
},
Error(String),
}
struct CachedEntry {
key: Zeroizing<[u8; 32]>,
kdf: KdfParams,
salt: [u8; 16],
}
struct State {
by_path: HashMap<String, CachedEntry>,
last_activity: Instant,
socket: PathBuf,
ttl: Duration,
}
type Shared = Arc<Mutex<State>>;
pub fn serve(socket: &Path, ttl: Duration) -> Result<()> {
if let Some(dir) = socket.parent() {
fs::DirBuilder::new()
.recursive(true)
.mode(0o700)
.create(dir)
.map_err(|e| {
crate::error::Error::Other(format!(
"agent: cannot create socket dir {}: {e}",
dir.display()
))
})?;
}
let _ = fs::remove_file(socket);
let listener = UnixListener::bind(socket).map_err(|e| {
crate::error::Error::Other(format!("agent: bind {}: {e}", socket.display()))
})?;
let _ = fs::set_permissions(socket, fs::Permissions::from_mode(0o600));
let state: Shared = Arc::new(Mutex::new(State {
by_path: HashMap::new(),
last_activity: Instant::now(),
socket: socket.to_path_buf(),
ttl,
}));
let watch_state = Arc::clone(&state);
thread::spawn(move || idle_watch(watch_state));
for stream in listener.incoming() {
match stream {
Ok(s) => {
let st = Arc::clone(&state);
thread::spawn(move || {
let _ = handle_conn(s, st);
});
}
Err(_) => continue,
}
}
Ok(())
}
fn idle_watch(state: Shared) {
loop {
thread::sleep(Duration::from_secs(1));
let (idle, ttl, sock) = {
let s = state.lock().unwrap();
(s.last_activity.elapsed(), s.ttl, s.socket.clone())
};
if idle >= ttl {
state.lock().unwrap().by_path.clear();
let _ = fs::remove_file(&sock);
std::process::exit(0);
}
}
}
fn handle_conn(stream: UnixStream, state: Shared) -> Result<()> {
let _ = stream.set_read_timeout(Some(Duration::from_secs(5)));
let mut reader = BufReader::new(&stream);
let mut writer = BufWriter::new(&stream);
let hs: Handshake = read_json(&mut reader)?;
if hs.proto != PROTOCOL_VERSION {
write_msg(&mut writer, &Response::Error("protocol mismatch".into()))?;
return Ok(());
}
let req: Request = read_json(&mut reader)?;
let resp = {
let mut s = state.lock().unwrap();
s.last_activity = Instant::now();
match req {
Request::Get { path } => match s.by_path.get(&path) {
Some(e) => Response::Got {
key: *e.key,
kdf: e.kdf,
salt: e.salt,
},
None => Response::Miss,
},
Request::Put {
path,
key,
kdf,
salt,
} => {
s.by_path.insert(
path,
CachedEntry {
key: Zeroizing::new(key),
kdf,
salt,
},
);
Response::Ok
}
Request::Forget { path } => {
s.by_path.remove(&path);
Response::Ok
}
Request::Status => {
let pid = std::process::id();
let socket = s.socket.to_string_lossy().into_owned();
let vaults = s.by_path.keys().cloned().collect();
let idle_secs = s.last_activity.elapsed().as_secs();
let ttl_secs = s.ttl.as_secs();
Response::StatusResp {
pid,
socket,
vaults,
idle_secs,
ttl_secs,
}
}
Request::Lock => {
s.by_path.clear();
Response::Ok
}
Request::Stop => {
let sock = s.socket.clone();
s.by_path.clear();
drop(s);
let _ = write_msg(&mut writer, &Response::Ok);
let _ = writer.flush();
let _ = fs::remove_file(&sock);
std::process::exit(0);
}
}
};
write_msg(&mut writer, &resp)?;
writer.flush()?;
Ok(())
}
pub fn try_get_key(path: &Path) -> Option<(MasterKey, KdfParams, [u8; 16])> {
if !super::enabled() {
return None;
}
match request(&Request::Get {
path: canonical(path),
})? {
Response::Got { key, kdf, salt } => Some((MasterKey::from_bytes(key), kdf, salt)),
_ => None,
}
}
pub fn put_key(path: &Path, key: &MasterKey, kdf: &KdfParams, salt: [u8; 16]) {
if !super::enabled() {
return;
}
let _ = request(&Request::Put {
path: canonical(path),
key: *key.as_bytes(),
kdf: *kdf,
salt,
});
}
pub fn forget(path: &Path) {
if !super::enabled() {
return;
}
let _ = request(&Request::Forget {
path: canonical(path),
});
}
pub fn lock_all() {
let _ = request(&Request::Lock);
}
pub fn stop() {
let _ = request(&Request::Stop);
}
pub fn status() -> Option<StatusInfo> {
match request(&Request::Status)? {
Response::StatusResp {
pid,
socket,
vaults,
idle_secs,
ttl_secs,
} => Some(StatusInfo {
pid,
socket,
vaults,
idle_secs,
ttl_secs,
}),
_ => None,
}
}
fn request(req: &Request) -> Option<Response> {
let stream = connect()?;
{
let mut w = BufWriter::new(&stream);
write_msg(&mut w, &Handshake { proto: PROTOCOL_VERSION }).ok()?;
write_msg(&mut w, req).ok()?;
w.flush().ok()?;
}
let mut r = BufReader::new(&stream);
read_json::<_, Response>(&mut r).ok()
}
fn connect() -> Option<UnixStream> {
let sock = socket_path()?;
if let Ok(s) = UnixStream::connect(&sock) {
return Some(s);
}
autostart(&sock);
for _ in 0..40 {
if let Ok(s) = UnixStream::connect(&sock) {
return Some(s);
}
thread::sleep(Duration::from_millis(50));
}
None
}
fn autostart(sock: &Path) {
let Ok(exe) = std::env::current_exe() else {
return;
};
let mut cmd = Command::new(exe);
cmd.arg("agent")
.arg("serve")
.arg("--socket")
.arg(sock)
.arg("--ttl")
.arg(super::ttl_secs().to_string());
cmd.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.process_group(0);
let _ = cmd.spawn();
}
fn socket_path() -> Option<PathBuf> {
if let Some(p) = std::env::var_os("ZKV_AGENT_SOCKET") {
return Some(PathBuf::from(p));
}
if let Some(dir) = std::env::var_os("XDG_RUNTIME_DIR") {
return Some(PathBuf::from(dir).join("zkv-agent.sock"));
}
Some(
std::env::temp_dir()
.join(format!("zkv-agent-{}", uid()))
.join("zkv-agent.sock"),
)
}
fn uid() -> u32 {
unsafe extern "C" {
fn getuid() -> u32;
}
unsafe { getuid() }
}
fn canonical(path: &Path) -> String {
fs::canonicalize(path)
.map(|p| p.to_string_lossy().into_owned())
.unwrap_or_else(|_| path.to_string_lossy().into_owned())
}
fn write_msg<W: Write>(w: &mut W, msg: &impl Serialize) -> Result<()> {
serde_json::to_writer(&mut *w, msg)?;
w.write_all(b"\n")?;
Ok(())
}
fn read_json<R: BufRead, T: DeserializeOwned>(r: &mut R) -> Result<T> {
let mut line = String::new();
r.read_line(&mut line)?;
Ok(serde_json::from_str(line.trim())?)
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::{Mutex, OnceLock};
fn env_lock() -> &'static Mutex<()> {
static M: OnceLock<Mutex<()>> = OnceLock::new();
M.get_or_init(|| Mutex::new(()))
}
fn unique_socket() -> PathBuf {
use std::sync::atomic::{AtomicU64, Ordering};
static N: AtomicU64 = AtomicU64::new(0);
let id = N.fetch_add(1, Ordering::SeqCst);
std::env::temp_dir().join(format!(
"zkv-agent-test-{}-{}.sock",
std::process::id(),
id
))
}
fn start_test_server(sock: &Path) -> Shared {
let _ = fs::remove_file(sock);
let listener = UnixListener::bind(sock).unwrap();
let state: Shared = Arc::new(Mutex::new(State {
by_path: HashMap::new(),
last_activity: Instant::now(),
socket: sock.to_path_buf(),
ttl: Duration::from_secs(600),
}));
let st = Arc::clone(&state);
thread::spawn(move || {
for s in listener.incoming() {
let Ok(s) = s else { continue };
let st2 = Arc::clone(&st);
thread::spawn(move || {
let _ = handle_conn(s, st2);
});
}
});
state
}
#[test]
fn serde_roundtrip_all_variants() {
let kdf = KdfParams {
m_kib: 1,
t_cost: 2,
p_cost: 3,
};
let reqs = vec![
Request::Get {
path: "/x.zkv".into(),
},
Request::Put {
path: "/x.zkv".into(),
key: [4u8; 32],
kdf,
salt: [9u8; 16],
},
Request::Forget {
path: "/y".into(),
},
Request::Status,
Request::Stop,
Request::Lock,
];
for r in reqs {
let s = serde_json::to_string(&r).unwrap();
let back: Request = serde_json::from_str(&s).unwrap();
assert_eq!(s, serde_json::to_string(&back).unwrap(), "Request not stable");
}
let resps = vec![
Response::Got {
key: [1u8; 32],
kdf,
salt: [2u8; 16],
},
Response::Miss,
Response::Ok,
Response::StatusResp {
pid: 42,
socket: "/tmp/s.sock".into(),
vaults: vec!["/a".into(), "/b".into()],
idle_secs: 7,
ttl_secs: 300,
},
Response::Error("boom".into()),
];
for r in resps {
let s = serde_json::to_string(&r).unwrap();
let back: Response = serde_json::from_str(&s).unwrap();
assert_eq!(s, serde_json::to_string(&back).unwrap(), "Response not stable");
}
}
#[test]
fn canonical_resolves_existing_file() {
let tmp = std::env::temp_dir().join(format!("zkv-canon-{}.bin", std::process::id()));
std::fs::write(&tmp, b"x").unwrap();
let with_dot = tmp.parent().unwrap().join(".").join(tmp.file_name().unwrap());
assert_eq!(canonical(&tmp), canonical(&with_dot));
let _ = std::fs::remove_file(&tmp);
}
#[test]
fn ttl_secs_parsing_mirrors_ui() {
let _g = env_lock().lock().unwrap();
unsafe {
std::env::remove_var("ZKV_LOCK_SECS");
assert_eq!(super::super::ttl_secs(), 300); std::env::set_var("ZKV_LOCK_SECS", "0");
assert_eq!(super::super::ttl_secs(), 0); std::env::set_var("ZKV_LOCK_SECS", "120");
assert_eq!(super::super::ttl_secs(), 120);
std::env::set_var("ZKV_LOCK_SECS", "garbage");
assert_eq!(super::super::ttl_secs(), 300); std::env::remove_var("ZKV_LOCK_SECS");
}
}
#[test]
fn enabled_gate() {
let _g = env_lock().lock().unwrap();
unsafe {
std::env::remove_var("ZKV_NO_AGENT");
std::env::set_var("ZKV_LOCK_SECS", "300");
assert!(super::super::enabled());
std::env::set_var("ZKV_NO_AGENT", "1");
assert!(!super::super::enabled()); std::env::remove_var("ZKV_NO_AGENT");
std::env::set_var("ZKV_LOCK_SECS", "0");
assert!(!super::super::enabled()); std::env::remove_var("ZKV_NO_AGENT");
std::env::remove_var("ZKV_LOCK_SECS");
}
}
#[test]
fn client_server_cache_roundtrip() {
let _g = env_lock().lock().unwrap();
let sock = unique_socket();
let _state = start_test_server(&sock);
unsafe {
std::env::set_var("ZKV_AGENT_SOCKET", &sock);
std::env::remove_var("ZKV_NO_AGENT");
std::env::set_var("ZKV_LOCK_SECS", "600");
}
thread::sleep(Duration::from_millis(50));
let vp = Path::new("/tmp/zkv-agent-no-such-vault.zkv");
assert!(try_get_key(vp).is_none());
let key = MasterKey::from_bytes([7u8; 32]);
let kdf = KdfParams::default();
let salt = [1u8; 16];
put_key(vp, &key, &kdf, salt);
let (gk, _gkdf, gsalt) = try_get_key(vp).expect("cached after put");
assert_eq!(gk.as_bytes(), &[7u8; 32]);
assert_eq!(gsalt, salt);
forget(vp);
assert!(try_get_key(vp).is_none());
unsafe {
std::env::remove_var("ZKV_AGENT_SOCKET");
std::env::remove_var("ZKV_LOCK_SECS");
}
let _ = fs::remove_file(&sock);
}
#[test]
fn version_mismatch_returns_error() {
let sock = unique_socket();
let _state = start_test_server(&sock);
thread::sleep(Duration::from_millis(50));
let s = UnixStream::connect(&sock).unwrap();
{
let mut w = BufWriter::new(&s);
write_msg(&mut w, &Handshake { proto: 999 }).unwrap();
write_msg(&mut w, &Request::Get { path: "x".into() }).unwrap();
w.flush().unwrap();
}
let mut r = BufReader::new(&s);
let resp: Response = read_json(&mut r).unwrap();
assert!(matches!(resp, Response::Error(_)), "expected Error, got {resp:?}");
let _ = fs::remove_file(&sock);
}
}
}
#[cfg(unix)]
pub use imp::{
forget, lock_all, put_key, serve, status, stop, try_get_key,
};
#[cfg(not(unix))]
mod stub {
use super::StatusInfo;
use crate::crypto::{KdfParams, MasterKey};
use crate::error::{Error, Result};
use std::path::Path;
use std::time::Duration;
pub fn try_get_key(_path: &Path) -> Option<(MasterKey, KdfParams, [u8; 16])> {
None
}
pub fn put_key(_path: &Path, _key: &MasterKey, _kdf: &KdfParams, _salt: [u8; 16]) {}
pub fn forget(_path: &Path) {}
pub fn lock_all() {}
pub fn stop() {}
pub fn status() -> Option<StatusInfo> {
None
}
pub fn serve(_socket: &Path, _ttl: Duration) -> Result<()> {
Err(Error::Other(
"agent is not supported on this platform (Unix only)".into(),
))
}
}
#[cfg(not(unix))]
pub use stub::{forget, lock_all, put_key, serve, status, stop, try_get_key};