use crate::approval::GrantChoice;
use crate::approval::ipc::ApprovalRequest;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::UnixStream;
const READ_TIMEOUT: Duration = Duration::from_secs(75);
const WRITE_TIMEOUT: Duration = Duration::from_secs(5);
const BACKOFF_START: Duration = Duration::from_millis(250);
const BACKOFF_MAX: Duration = Duration::from_secs(5);
const MAX_LINE: usize = 64 * 1024;
#[derive(Debug, Clone, PartialEq)]
pub enum CompanionEvent {
Connected { socket: PathBuf },
Waiting,
Request { request: ApprovalRequest },
Acked { id: String, choice: GrantChoice },
Stale { id: String },
BadChoice { id: String },
Empty,
Disconnected { reason: String },
}
#[derive(Debug)]
pub struct CompanionChoice {
pub id: String,
pub choice: GrantChoice,
}
#[derive(Default)]
pub struct CompanionOptions {
pub once: bool,
pub request_id: Option<String>,
pub discover: bool,
}
pub struct EventSender {
tx: std::sync::mpsc::Sender<CompanionEvent>,
wake: Arc<dyn Fn() + Send + Sync>,
}
impl EventSender {
pub fn new(
tx: std::sync::mpsc::Sender<CompanionEvent>,
wake: impl Fn() + Send + Sync + 'static,
) -> Self {
Self {
tx,
wake: Arc::new(wake),
}
}
fn send(&self, event: CompanionEvent) -> Option<()> {
self.tx.send(event).ok()?;
(self.wake)();
Some(())
}
}
pub struct CompanionHandle {
pub choices: tokio::sync::mpsc::Sender<CompanionChoice>,
pub events: std::sync::mpsc::Receiver<CompanionEvent>,
}
pub fn spawn_companion(socket: PathBuf) -> CompanionHandle {
let (event_tx, event_rx) = std::sync::mpsc::channel::<CompanionEvent>();
let (choice_tx, choice_rx) = tokio::sync::mpsc::channel(8);
tokio::spawn(companion_loop(
socket,
EventSender::new(event_tx, || {}),
choice_rx,
CompanionOptions::default(),
));
CompanionHandle {
choices: choice_tx,
events: event_rx,
}
}
pub async fn companion_loop(
mut socket: PathBuf,
event_tx: EventSender,
mut choice_rx: tokio::sync::mpsc::Receiver<CompanionChoice>,
options: CompanionOptions,
) {
let mut backoff = BACKOFF_START;
loop {
let outcome = run_one_round(&socket, &event_tx, &mut choice_rx, &options).await;
if options.once || choice_rx.is_closed() {
break;
}
match outcome {
RoundOutcome::SocketAlive => backoff = BACKOFF_START,
RoundOutcome::SocketGone => {
tokio::time::sleep(backoff).await;
if options.discover {
socket = crate::approval::ipc::ApprovalIpc::default_socket_path();
}
backoff = std::cmp::min(backoff * 2, BACKOFF_MAX);
}
RoundOutcome::UiGone => break,
}
}
}
enum RoundOutcome {
SocketAlive,
SocketGone,
UiGone,
}
async fn run_one_round(
socket: &std::path::Path,
event_tx: &EventSender,
choice_rx: &mut tokio::sync::mpsc::Receiver<CompanionChoice>,
options: &CompanionOptions,
) -> RoundOutcome {
let stream = match UnixStream::connect(socket).await {
Ok(stream) => stream,
Err(e) => {
let _ = event_tx.send(CompanionEvent::Disconnected {
reason: format!("connect: {e}"),
});
return RoundOutcome::SocketGone;
}
};
let _ = event_tx.send(CompanionEvent::Connected {
socket: socket.to_path_buf(),
});
let (rd, mut wr) = stream.into_split();
let mut reader = BufReader::new(rd);
let hello = serde_json::json!({"op": "wait", "requestId": options.request_id});
if write_line(&mut wr, &hello).await.is_err() {
let _ = event_tx.send(CompanionEvent::Disconnected {
reason: "wait write failed".into(),
});
return RoundOutcome::SocketGone;
}
let _ = event_tx.send(CompanionEvent::Waiting);
let line = match tokio::time::timeout(READ_TIMEOUT, read_line(&mut reader)).await {
Ok(Some(line)) => line,
Ok(None) | Err(_) => {
let _ = event_tx.send(CompanionEvent::Disconnected {
reason: "connection closed or timed out while waiting".into(),
});
return RoundOutcome::SocketGone;
}
};
let Ok(msg) = serde_json::from_str::<serde_json::Value>(&line) else {
let _ = event_tx.send(CompanionEvent::Disconnected {
reason: "malformed line from server".into(),
});
return RoundOutcome::SocketGone;
};
match msg["op"].as_str().unwrap_or("") {
"empty" => {
let _ = event_tx.send(CompanionEvent::Empty);
return RoundOutcome::SocketAlive;
}
"request" => {}
other => {
let _ = event_tx.send(CompanionEvent::Disconnected {
reason: format!("unexpected op {other:?}"),
});
return RoundOutcome::SocketGone;
}
}
let Ok(request) = serde_json::from_value::<ApprovalRequest>(msg["request"].clone()) else {
let _ = event_tx.send(CompanionEvent::Disconnected {
reason: "unparseable request payload".into(),
});
return RoundOutcome::SocketGone;
};
let id = request.id.clone();
if options
.request_id
.as_ref()
.is_some_and(|expected| expected != &id)
{
let _ = event_tx.send(CompanionEvent::Stale {
id: options.request_id.clone().unwrap(),
});
return RoundOutcome::SocketAlive;
}
let _ = event_tx.send(CompanionEvent::Request { request });
let choice = loop {
tokio::select! {
biased;
_ = read_line(&mut reader) => {
let _ = event_tx.send(CompanionEvent::Stale { id });
return RoundOutcome::SocketGone;
}
choice = choice_rx.recv() => {
match choice {
Some(answer) if answer.id == id => break answer.choice,
Some(_) => continue,
None => return RoundOutcome::UiGone,
}
}
}
};
let choice_str = match choice {
GrantChoice::Once => "once",
GrantChoice::Session => "session",
GrantChoice::Decline => "decline",
};
let reply = serde_json::json!({"op": "reply", "id": id, "choice": choice_str});
if write_line(&mut wr, &reply).await.is_err() {
let _ = event_tx.send(CompanionEvent::Disconnected {
reason: "reply write failed".into(),
});
return RoundOutcome::SocketGone;
}
let ack_line = match tokio::time::timeout(READ_TIMEOUT, read_line(&mut reader)).await {
Ok(Some(line)) => line,
_ => {
let _ = event_tx.send(CompanionEvent::Disconnected {
reason: "connection closed before ack".into(),
});
return RoundOutcome::SocketGone;
}
};
let Ok(ack) = serde_json::from_str::<serde_json::Value>(&ack_line) else {
let _ = event_tx.send(CompanionEvent::Disconnected {
reason: "malformed ack".into(),
});
return RoundOutcome::SocketGone;
};
match ack["op"].as_str().unwrap_or("") {
"ok" => {
let _ = event_tx.send(CompanionEvent::Acked { id, choice });
}
"stale" => {
let _ = event_tx.send(CompanionEvent::Stale { id });
}
_ => {
let _ = event_tx.send(CompanionEvent::BadChoice { id });
}
}
RoundOutcome::SocketAlive
}
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<()> {
tokio::time::timeout(WRITE_TIMEOUT, async {
stream
.write_all(serde_json::to_string(value).unwrap_or_default().as_bytes())
.await?;
stream.write_all(b"\n").await?;
stream.flush().await
})
.await
.map_err(|_| std::io::Error::new(std::io::ErrorKind::TimedOut, "write timed out"))?
}
#[cfg(test)]
mod tests {
use super::*;
use crate::approval::ConfirmOutcome;
use crate::approval::ipc::ApprovalIpc;
use std::sync::Arc;
#[test]
fn event_delivery_wakes_the_ui_immediately() {
let (tx, rx) = std::sync::mpsc::channel();
let wakes = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let count = Arc::clone(&wakes);
let sender = EventSender::new(tx, move || {
count.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
});
sender.send(CompanionEvent::Waiting).unwrap();
assert_eq!(rx.try_recv().unwrap(), CompanionEvent::Waiting);
assert_eq!(wakes.load(std::sync::atomic::Ordering::SeqCst), 1);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn stale_queued_click_cannot_approve_a_new_request() {
let dir = tempfile::TempDir::new().unwrap();
let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
let handle = spawn_companion(ipc.socket_path().to_path_buf());
handle
.choices
.send(CompanionChoice {
id: "old".into(),
choice: GrantChoice::Once,
})
.await
.unwrap();
let ask = tokio::spawn({
let ipc = Arc::clone(&ipc);
async move { ipc.ask(sample_request("new")).await }
});
next_event(&handle.events, |e| {
matches!(e, CompanionEvent::Request { .. })
})
.await;
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(!ask.is_finished(), "stale click authorized the new request");
handle
.choices
.send(CompanionChoice {
id: "new".into(),
choice: GrantChoice::Decline,
})
.await
.unwrap();
assert_eq!(
ask.await.unwrap(),
ConfirmOutcome::Chosen(GrantChoice::Decline)
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn targeted_prompt_never_displays_a_different_request() {
let dir = tempfile::TempDir::new().unwrap();
let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
let ask = tokio::spawn({
let ipc = Arc::clone(&ipc);
async move { ipc.ask(sample_request("new")).await }
});
let (tx, rx) = std::sync::mpsc::channel();
let (_choices, answers) = tokio::sync::mpsc::channel(8);
tokio::time::sleep(Duration::from_millis(30)).await;
companion_loop(
ipc.socket_path().to_path_buf(),
EventSender::new(tx, || {}),
answers,
CompanionOptions {
once: true,
request_id: Some("old".into()),
discover: false,
},
)
.await;
let events: Vec<_> = rx.try_iter().collect();
assert!(events.contains(&CompanionEvent::Empty));
assert!(
!events
.iter()
.any(|e| matches!(e, CompanionEvent::Request { .. }))
);
assert!(!ask.is_finished());
ask.abort();
}
fn sample_request(id: &str) -> ApprovalRequest {
ApprovalRequest {
id: id.into(),
category: "write".into(),
connection: "local-dev".into(),
database: Some("app".into()),
tables: vec!["app.users".into()],
snippet: "UPDATE users SET name = 'x' WHERE id = 2".into(),
}
}
async fn next_event(
rx: &std::sync::mpsc::Receiver<CompanionEvent>,
pred: impl Fn(&CompanionEvent) -> bool,
) -> CompanionEvent {
let deadline = Instant::now() + Duration::from_secs(10);
loop {
while let Ok(event) = rx.try_recv() {
if pred(&event) {
return event;
}
}
assert!(Instant::now() < deadline, "timed out waiting for event");
tokio::time::sleep(Duration::from_millis(10)).await;
}
}
use std::time::Instant;
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn companion_round_trip_once() {
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 = spawn_companion(socket);
let ask = tokio::spawn({
let ipc = Arc::clone(&ipc);
async move { ipc.ask(sample_request("r1")).await }
});
let event = next_event(&handle.events, |e| {
matches!(e, CompanionEvent::Request { .. })
})
.await;
let CompanionEvent::Request { request } = event else {
unreachable!()
};
assert_eq!(request.id, "r1");
assert_eq!(request.tables, vec!["app.users".to_string()]);
handle
.choices
.send(CompanionChoice {
id: "r1".into(),
choice: GrantChoice::Once,
})
.await
.unwrap();
let ack = next_event(&handle.events, |e| {
matches!(e, CompanionEvent::Acked { .. })
})
.await;
assert_eq!(
ack,
CompanionEvent::Acked {
id: "r1".into(),
choice: GrantChoice::Once
}
);
assert_eq!(
ask.await.unwrap(),
ConfirmOutcome::Chosen(GrantChoice::Once)
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn companion_round_trip_decline() {
let dir = tempfile::TempDir::new().unwrap();
let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
let handle = spawn_companion(ipc.socket_path().to_path_buf());
let ask = tokio::spawn({
let ipc = Arc::clone(&ipc);
async move { ipc.ask(sample_request("r2")).await }
});
next_event(&handle.events, |e| {
matches!(e, CompanionEvent::Request { .. })
})
.await;
handle
.choices
.send(CompanionChoice {
id: "r2".into(),
choice: GrantChoice::Decline,
})
.await
.unwrap();
next_event(&handle.events, |e| {
matches!(e, CompanionEvent::Acked { .. })
})
.await;
assert_eq!(
ask.await.unwrap(),
ConfirmOutcome::Chosen(GrantChoice::Decline)
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn companion_survives_server_restart() {
let dir = tempfile::TempDir::new().unwrap();
let socket = {
let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
ipc.socket_path().to_path_buf()
};
let handle = spawn_companion(socket.clone());
let gone = next_event(&handle.events, |e| {
matches!(e, CompanionEvent::Disconnected { .. })
})
.await;
let CompanionEvent::Disconnected { reason } = gone else {
unreachable!()
};
assert!(!reason.is_empty());
let ipc2 = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
assert_eq!(ipc2.socket_path(), socket.as_path());
let ask = tokio::spawn({
let ipc = Arc::clone(&ipc2);
async move { ipc.ask(sample_request("r3")).await }
});
next_event(&handle.events, |e| {
matches!(e, CompanionEvent::Request { .. })
})
.await;
handle
.choices
.send(CompanionChoice {
id: "r3".into(),
choice: GrantChoice::Session,
})
.await
.unwrap();
next_event(&handle.events, |e| {
matches!(e, CompanionEvent::Acked { .. })
})
.await;
assert_eq!(
ask.await.unwrap(),
ConfirmOutcome::Chosen(GrantChoice::Session)
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn companion_reports_missing_server() {
let dir = tempfile::TempDir::new().unwrap();
let handle = spawn_companion(dir.path().join("nope.sock"));
next_event(&handle.events, |e| {
matches!(e, CompanionEvent::Disconnected { .. })
})
.await;
next_event(&handle.events, |e| {
matches!(e, CompanionEvent::Disconnected { .. })
})
.await;
}
}