use crate::approval::{ConfirmOutcome, GrantChoice};
use std::os::unix::fs::PermissionsExt;
use std::path::PathBuf;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::UnixStream;
pub const APPROVAL_IPC_TIMEOUT: Duration = Duration::from_secs(60);
const MAX_LINE: usize = 64 * 1024;
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct ApprovalRequest {
pub id: String,
pub category: String,
pub connection: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub database: Option<String>,
pub tables: Vec<String>,
pub snippet: String,
}
struct Slot {
request: ApprovalRequest,
tx: tokio::sync::oneshot::Sender<GrantChoice>,
deadline: Instant,
}
struct PendingGuard {
inner: Arc<Mutex<Inner>>,
id: String,
}
impl Drop for PendingGuard {
fn drop(&mut self) {
let mut inner = self.inner.lock().unwrap();
if inner.slot.as_ref().is_some_and(|s| s.request.id == self.id) {
inner.slot = None;
}
}
}
struct Inner {
slot: Option<Slot>,
}
pub struct ApprovalIpc {
inner: Arc<Mutex<Inner>>,
socket_path: PathBuf,
session_file: Option<PathBuf>,
worker: tokio::task::JoinHandle<()>,
started: Instant,
served: AtomicU64,
launch_gui: bool,
}
impl ApprovalIpc {
pub fn start() -> std::io::Result<Arc<Self>> {
Self::start_with_gui(true)
}
pub fn start_with_gui(launch_gui: bool) -> std::io::Result<Arc<Self>> {
Self::bind_at(
crate::app::paths::runtime_dir(),
&format!("p{}.sock", std::process::id()),
launch_gui,
)
}
pub fn start_at(dir: Option<PathBuf>) -> std::io::Result<Arc<Self>> {
Self::bind_at(
dir.unwrap_or_else(crate::app::paths::runtime_dir),
"approval.sock",
false,
)
}
fn bind_at(runtime: PathBuf, filename: &str, launch_gui: bool) -> std::io::Result<Arc<Self>> {
std::fs::create_dir_all(&runtime)?;
std::fs::set_permissions(&runtime, std::fs::Permissions::from_mode(0o700))?;
let socket_path = runtime.join(filename);
if std::os::unix::net::UnixStream::connect(&socket_path).is_ok() {
return Err(std::io::Error::new(
std::io::ErrorKind::AddrInUse,
"approval socket is already serving",
));
}
let _ = std::fs::remove_file(&socket_path);
let listener = tokio::net::UnixListener::bind(&socket_path)?;
std::fs::set_permissions(&socket_path, std::fs::Permissions::from_mode(0o600))?;
let session_file = runtime.join("sessions").join(format!("{filename}.json"));
if let Some(parent) = session_file.parent() {
std::fs::create_dir_all(parent)?;
}
let registry = serde_json::json!({
"pid": std::process::id(),
"socket": socket_path.display().to_string(),
"started": iso_now(),
});
let _ = std::fs::write(
&session_file,
serde_json::to_string(®istry).unwrap_or_default(),
);
let inner = Arc::new(Mutex::new(Inner { slot: None }));
let accept_inner = Arc::clone(&inner);
let worker = tokio::spawn(async move {
loop {
let Ok((stream, _)) = listener.accept().await else {
break;
};
let inner = Arc::clone(&accept_inner);
tokio::spawn(handle_connection(stream, inner));
}
});
Ok(Arc::new(Self {
inner,
socket_path,
session_file: Some(session_file),
worker,
started: Instant::now(),
served: AtomicU64::new(0),
launch_gui,
}))
}
pub fn default_socket_path() -> PathBuf {
if let Some(session) = crate::gui::live_sessions()
.into_iter()
.find(|s| s.alive && std::path::Path::new(&s.socket).exists())
{
return PathBuf::from(session.socket);
}
crate::app::paths::runtime_dir().join("approval.sock")
}
pub fn socket_path(&self) -> &std::path::Path {
&self.socket_path
}
pub fn answered_count(&self) -> u64 {
self.served.load(Ordering::Relaxed)
}
pub fn uptime(&self) -> Duration {
self.started.elapsed()
}
pub async fn ask(&self, request: ApprovalRequest) -> ConfirmOutcome {
self.ask_with_timeout(request, APPROVAL_IPC_TIMEOUT).await
}
async fn ask_with_timeout(
&self,
request: ApprovalRequest,
timeout: Duration,
) -> ConfirmOutcome {
let id = request.id.clone();
let (tx, rx) = tokio::sync::oneshot::channel::<GrantChoice>();
{
let mut inner = self.inner.lock().unwrap();
if inner.slot.is_some() {
return ConfirmOutcome::Unavailable {
reason: "another approval request is already pending".into(),
};
}
inner.slot = Some(Slot {
request,
tx,
deadline: Instant::now() + timeout,
});
}
let _pending = PendingGuard {
inner: Arc::clone(&self.inner),
id: id.clone(),
};
let mut dialog = if self.launch_gui {
crate::gui::launch::open_prompt(&self.socket_path, &id)
} else {
None
};
let answer = async {
if let Some(child) = dialog.as_mut() {
tokio::select! {
biased;
result = rx => result.ok(),
_ = child.wait() => None,
}
} else {
rx.await.ok()
}
};
let outcome = match tokio::time::timeout(timeout, answer).await {
Ok(Some(choice)) => {
self.served.fetch_add(1, Ordering::Relaxed);
ConfirmOutcome::Chosen(choice)
}
Ok(None) => ConfirmOutcome::Unavailable {
reason: "approval window closed or companion dropped the request without a choice"
.into(),
},
Err(_elapsed) => {
ConfirmOutcome::Unavailable {
reason: format!("approval IPC timed out after {}s", timeout.as_secs()),
}
}
};
if let Some(child) = dialog.as_mut() {
let _ = tokio::time::timeout(Duration::from_millis(500), child.wait()).await;
}
outcome
}
}
impl Drop for ApprovalIpc {
fn drop(&mut self) {
self.inner.lock().unwrap().slot = None;
self.worker.abort();
let _ = std::fs::remove_file(&self.socket_path);
if let Some(session) = &self.session_file {
let _ = std::fs::remove_file(session);
}
}
}
fn peer_is_same_uid<Fd: std::os::unix::io::AsRawFd>(stream: &Fd) -> bool {
#[cfg(target_os = "macos")]
{
let mut cred: libc::xucred = unsafe { std::mem::zeroed() };
let mut len = std::mem::size_of_val(&cred) as libc::socklen_t;
let rc = unsafe {
libc::getsockopt(
stream.as_raw_fd(),
libc::SOL_LOCAL,
libc::LOCAL_PEERCRED,
&mut cred as *mut _ as *mut libc::c_void,
&mut len,
)
};
rc == 0 && cred.cr_uid == unsafe { libc::geteuid() }
}
#[cfg(target_os = "linux")]
{
let mut cred: libc::ucred = unsafe { std::mem::zeroed() };
let mut len = std::mem::size_of_val(&cred) as libc::socklen_t;
let rc = unsafe {
libc::getsockopt(
stream.as_raw_fd(),
libc::SOL_SOCKET,
libc::SO_PEERCRED,
&mut cred as *mut _ as *mut libc::c_void,
&mut len,
)
};
rc == 0 && cred.uid == unsafe { libc::geteuid() }
}
#[cfg(not(any(target_os = "macos", target_os = "linux")))]
{
let _ = stream;
false
}
}
async fn handle_connection(stream: UnixStream, inner: Arc<Mutex<Inner>>) {
if !peer_is_same_uid(&stream) {
return;
}
let (rd, mut wr) = stream.into_split();
let mut reader = BufReader::new(rd);
let Some(first) = read_line(&mut reader).await else {
return;
};
let Ok(cmd) = serde_json::from_str::<serde_json::Value>(&first) else {
return;
};
if cmd["op"].as_str() != Some("wait") {
return;
}
let deadline = Instant::now() + APPROVAL_IPC_TIMEOUT;
loop {
{
let guard = inner.lock().unwrap();
if guard.slot.is_some() || cmd["requestId"].is_string() || Instant::now() >= deadline {
break;
}
}
tokio::time::sleep(Duration::from_millis(25)).await;
}
let request = inner
.lock()
.unwrap()
.slot
.as_ref()
.filter(|slot| {
cmd["requestId"]
.as_str()
.is_none_or(|id| id == slot.request.id)
})
.map(|slot| serde_json::to_value(&slot.request).unwrap_or_default());
let Some(request) = request else {
let _ = write_line(&mut wr, &serde_json::json!({"op": "empty"})).await;
return;
};
if write_line(
&mut wr,
&serde_json::json!({
"op": "request",
"request": request,
}),
)
.await
.is_err()
{
return;
}
let request_id = request["id"].as_str().unwrap_or_default().to_string();
let expired = async {
loop {
let active = inner.lock().unwrap().slot.as_ref().is_some_and(|s| {
s.request.id == request_id && Instant::now() < s.deadline && !s.tx.is_closed()
});
if !active {
break;
}
tokio::time::sleep(Duration::from_millis(25)).await;
}
};
let line = tokio::select! {
biased;
_ = expired => {
let _ = write_line(&mut wr, &serde_json::json!({"op": "stale"})).await;
return;
}
line = read_line(&mut reader) => line,
};
let Some(line) = line else {
return;
};
let Ok(reply) = serde_json::from_str::<serde_json::Value>(&line) else {
return;
};
if reply["op"].as_str() != Some("reply") {
return;
}
enum Reply {
Choice(GrantChoice, Slot),
Stale,
BadChoice,
}
let decision = {
let mut guard = inner.lock().unwrap();
let id_matches = guard
.slot
.as_ref()
.map(|s| {
s.request.id == request_id
&& s.request.id == reply["id"].as_str().unwrap_or("")
&& Instant::now() < s.deadline
&& !s.tx.is_closed()
})
.unwrap_or(false);
if !id_matches {
Reply::Stale
} else {
let choice = match reply["choice"].as_str() {
Some("once") => Some(GrantChoice::Once),
Some("session") => Some(GrantChoice::Session),
Some("decline") => Some(GrantChoice::Decline),
_ => None,
};
match choice {
Some(choice) => Reply::Choice(choice, guard.slot.take().unwrap()),
None => Reply::BadChoice,
}
}
};
match decision {
Reply::Stale => {
let _ = write_line(&mut wr, &serde_json::json!({"op": "stale"})).await;
}
Reply::BadChoice => {
let _ = write_line(&mut wr, &serde_json::json!({"op": "bad-choice"})).await;
}
Reply::Choice(choice, slot) => {
if write_line(&mut wr, &serde_json::json!({"op": "ok"}))
.await
.is_ok()
{
let _ = slot.tx.send(choice);
}
}
}
}
async fn read_line<R: tokio::io::AsyncBufRead + Unpin>(reader: &mut R) -> Option<String> {
let mut line = String::new();
reader.read_line(&mut line).await.ok()?;
let line = line.trim().to_string();
if line.is_empty() || line.len() > MAX_LINE {
return None;
}
Some(line)
}
async fn write_line<W: tokio::io::AsyncWrite + Unpin>(
stream: &mut W,
value: &serde_json::Value,
) -> std::io::Result<()> {
stream
.write_all(serde_json::to_string(value).unwrap_or_default().as_bytes())
.await?;
stream.write_all(b"\n").await?;
stream.flush().await
}
fn iso_now() -> String {
time::OffsetDateTime::now_utc()
.format(&time::format_description::well_known::Rfc3339)
.unwrap_or_else(|_| "1970-01-01T00:00:00Z".into())
}
#[cfg(test)]
mod tests {
use super::*;
fn request(id: &str) -> ApprovalRequest {
ApprovalRequest {
id: id.into(),
category: "write".into(),
connection: "local-dev".into(),
database: Some("app".into()),
tables: vec![],
snippet: "UPDATE items SET id = 2".into(),
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn independent_server_sockets_do_not_replace_each_other() {
let dir = tempfile::TempDir::new().unwrap();
let first = ApprovalIpc::bind_at(dir.path().to_path_buf(), "p1.sock", false).unwrap();
let second = ApprovalIpc::bind_at(dir.path().to_path_buf(), "p2.sock", false).unwrap();
assert!(ApprovalIpc::bind_at(dir.path().to_path_buf(), "p1.sock", false).is_err());
let answer_first = tokio::spawn(approver_script(
first.socket_path().to_path_buf(),
"decline",
));
let answer_second =
tokio::spawn(approver_script(second.socket_path().to_path_buf(), "once"));
let (a, b) = tokio::join!(first.ask(request("a")), second.ask(request("b")));
assert_eq!(a, ConfirmOutcome::Chosen(GrantChoice::Decline));
assert_eq!(b, ConfirmOutcome::Chosen(GrantChoice::Once));
answer_first.await.unwrap().unwrap();
answer_second.await.unwrap().unwrap();
drop(first);
assert!(second.socket_path().exists());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn timeout_notifies_idle_dialog_and_cancel_releases_slot() {
let dir = tempfile::TempDir::new().unwrap();
let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
let handle = crate::gui::companion::spawn_companion(ipc.socket_path().to_path_buf());
let result = ipc
.ask_with_timeout(request("expired"), Duration::from_millis(150))
.await;
assert!(matches!(result, ConfirmOutcome::Unavailable { .. }));
tokio::time::timeout(Duration::from_secs(2), async {
loop {
if handle
.events
.try_iter()
.any(|e| matches!(e, crate::gui::companion::CompanionEvent::Stale { .. }))
{
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.unwrap();
assert_eq!(ipc.answered_count(), 0);
let ask = tokio::spawn({
let ipc = Arc::clone(&ipc);
async move { ipc.ask(request("cancelled")).await }
});
tokio::time::sleep(Duration::from_millis(30)).await;
ask.abort();
let _ = ask.await;
assert!(ipc.inner.lock().unwrap().slot.is_none());
}
async fn approver_script(socket: PathBuf, choice: &'static str) -> Option<()> {
let stream = UnixStream::connect(socket).await.ok()?;
let (rd, mut wr) = stream.into_split();
let mut reader = BufReader::new(rd);
let hello = serde_json::json!({"op": "wait"});
write_line(&mut wr, &hello).await.ok()?;
let line = read_line(&mut reader).await?;
let msg: serde_json::Value = serde_json::from_str(&line).ok()?;
assert_eq!(msg["op"], "request", "{msg}");
let id = msg["request"]["id"].as_str()?.to_string();
let reply = serde_json::json!({"op": "reply", "id": id, "choice": choice});
write_line(&mut wr, &reply).await.ok()?;
let ack = read_line(&mut reader).await?;
let ack: serde_json::Value = serde_json::from_str(&ack).ok()?;
assert_eq!(ack["op"], "ok", "{ack}");
Some(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn approve_once_round_trip() {
let dir = tempfile::TempDir::new().unwrap();
let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
let socket = ipc.socket_path().to_path_buf();
let handle = tokio::spawn(approver_script(socket, "once"));
let outcome = ipc
.ask(ApprovalRequest {
id: uuid::Uuid::new_v4().to_string(),
category: "write".into(),
connection: "c".into(),
database: Some("app".into()),
tables: vec!["app.users".into()],
snippet: "UPDATE users SET id = 2".into(),
})
.await;
assert_eq!(outcome, ConfirmOutcome::Chosen(GrantChoice::Once));
assert!(handle.await.unwrap().is_some());
assert_eq!(ipc.answered_count(), 1);
assert!(ipc.socket_path().exists());
drop(ipc);
assert!(!ipc_exists(dir.path()), "socket removed on drop");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn decline_and_session_choices() {
let dir = tempfile::TempDir::new().unwrap();
let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
for choice in ["decline", "session"] {
let socket = ipc.socket_path().to_path_buf();
let leaked: &'static str = Box::leak(choice.to_string().into_boxed_str());
let handle = tokio::spawn(approver_script(socket, leaked));
let outcome = ipc
.ask(ApprovalRequest {
id: uuid::Uuid::new_v4().to_string(),
category: "write".into(),
connection: "c".into(),
database: None,
tables: vec![],
snippet: "s".into(),
})
.await;
let expected = if choice == "decline" {
GrantChoice::Decline
} else {
GrantChoice::Session
};
assert_eq!(outcome, ConfirmOutcome::Chosen(expected));
assert!(handle.await.unwrap().is_some());
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn no_companion_fails_closed_after_deadline() {
let dir = tempfile::TempDir::new().unwrap();
let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
let ipc2 = Arc::clone(&ipc);
let first = tokio::spawn(async move {
ipc2.ask(ApprovalRequest {
id: "first".into(),
category: "write".into(),
connection: "c".into(),
database: None,
tables: vec![],
snippet: "s".into(),
})
.await
});
tokio::time::sleep(Duration::from_millis(50)).await;
let second = ipc
.ask(ApprovalRequest {
id: "second".into(),
category: "write".into(),
connection: "c".into(),
database: None,
tables: vec![],
snippet: "s".into(),
})
.await;
assert!(
matches!(second, ConfirmOutcome::Unavailable { .. }),
"second concurrent ask fails closed: {second:?}"
);
first.abort();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn wrong_id_reply_is_stale() {
let dir = tempfile::TempDir::new().unwrap();
let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
let socket = ipc.socket_path().to_path_buf();
let ask = tokio::spawn({
let ipc = Arc::clone(&ipc);
async move {
ipc.ask(ApprovalRequest {
id: "real-id".into(),
category: "write".into(),
connection: "c".into(),
database: None,
tables: vec![],
snippet: "s".into(),
})
.await
}
});
tokio::time::sleep(Duration::from_millis(100)).await;
let stream = UnixStream::connect(&socket).await.unwrap();
let (rd, mut wr) = stream.into_split();
let mut reader = BufReader::new(rd);
write_line(&mut wr, &serde_json::json!({"op": "wait"}))
.await
.unwrap();
let line = read_line(&mut reader).await.unwrap();
let msg: serde_json::Value = serde_json::from_str(&line).unwrap();
assert_eq!(msg["op"], "request");
write_line(
&mut wr,
&serde_json::json!({"op": "reply", "id": "forged", "choice": "once"}),
)
.await
.unwrap();
let ack = read_line(&mut reader).await.unwrap();
let ack: serde_json::Value = serde_json::from_str(&ack).unwrap();
assert_eq!(ack["op"], "stale", "{ack}");
ask.abort();
}
fn ipc_exists(dir: &std::path::Path) -> bool {
dir.join("approval.sock").exists()
}
}