1use crate::approval::{ConfirmOutcome, GrantChoice};
18use std::os::unix::fs::PermissionsExt;
19use std::path::PathBuf;
20use std::sync::atomic::{AtomicU64, Ordering};
21use std::sync::{Arc, Mutex};
22use std::time::{Duration, Instant};
23use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
24use tokio::net::UnixStream;
25
26pub const APPROVAL_IPC_TIMEOUT: Duration = Duration::from_secs(60);
28
29const MAX_LINE: usize = 64 * 1024;
31
32#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
33pub struct ApprovalRequest {
34 pub id: String,
35 pub category: String,
36 pub connection: String,
37 #[serde(default, skip_serializing_if = "Option::is_none")]
38 pub database: Option<String>,
39 pub tables: Vec<String>,
40 pub snippet: String,
41}
42
43struct Slot {
44 request: ApprovalRequest,
45 tx: tokio::sync::oneshot::Sender<GrantChoice>,
46 deadline: Instant,
47}
48
49struct PendingGuard {
51 inner: Arc<Mutex<Inner>>,
52 id: String,
53}
54
55impl Drop for PendingGuard {
56 fn drop(&mut self) {
57 let mut inner = self.inner.lock().unwrap();
58 if inner.slot.as_ref().is_some_and(|s| s.request.id == self.id) {
59 inner.slot = None;
60 }
61 }
62}
63
64struct Inner {
65 slot: Option<Slot>,
66}
67
68pub struct ApprovalIpc {
69 inner: Arc<Mutex<Inner>>,
70 socket_path: PathBuf,
71 session_file: Option<PathBuf>,
72 worker: tokio::task::JoinHandle<()>,
73 started: Instant,
74 served: AtomicU64,
75 launch_gui: bool,
76}
77
78impl ApprovalIpc {
79 pub fn start() -> std::io::Result<Arc<Self>> {
83 Self::start_with_gui(true)
84 }
85
86 pub fn start_with_gui(launch_gui: bool) -> std::io::Result<Arc<Self>> {
87 Self::bind_at(
90 crate::app::paths::runtime_dir(),
91 &format!("p{}.sock", std::process::id()),
92 launch_gui,
93 )
94 }
95
96 pub fn start_at(dir: Option<PathBuf>) -> std::io::Result<Arc<Self>> {
99 Self::bind_at(
100 dir.unwrap_or_else(crate::app::paths::runtime_dir),
101 "approval.sock",
102 false,
103 )
104 }
105
106 fn bind_at(runtime: PathBuf, filename: &str, launch_gui: bool) -> std::io::Result<Arc<Self>> {
107 std::fs::create_dir_all(&runtime)?;
108 std::fs::set_permissions(&runtime, std::fs::Permissions::from_mode(0o700))?;
109 let socket_path = runtime.join(filename);
110 if std::os::unix::net::UnixStream::connect(&socket_path).is_ok() {
111 return Err(std::io::Error::new(
112 std::io::ErrorKind::AddrInUse,
113 "approval socket is already serving",
114 ));
115 }
116 let _ = std::fs::remove_file(&socket_path);
117 let listener = tokio::net::UnixListener::bind(&socket_path)?;
118 std::fs::set_permissions(&socket_path, std::fs::Permissions::from_mode(0o600))?;
119
120 let session_file = runtime.join("sessions").join(format!("{filename}.json"));
122 if let Some(parent) = session_file.parent() {
123 std::fs::create_dir_all(parent)?;
124 }
125 let registry = serde_json::json!({
126 "pid": std::process::id(),
127 "socket": socket_path.display().to_string(),
128 "started": iso_now(),
129 });
130 let _ = std::fs::write(
131 &session_file,
132 serde_json::to_string(®istry).unwrap_or_default(),
133 );
134
135 let inner = Arc::new(Mutex::new(Inner { slot: None }));
136 let accept_inner = Arc::clone(&inner);
137 let worker = tokio::spawn(async move {
138 loop {
139 let Ok((stream, _)) = listener.accept().await else {
140 break;
141 };
142 let inner = Arc::clone(&accept_inner);
143 tokio::spawn(handle_connection(stream, inner));
144 }
145 });
146
147 Ok(Arc::new(Self {
148 inner,
149 socket_path,
150 session_file: Some(session_file),
151 worker,
152 started: Instant::now(),
153 served: AtomicU64::new(0),
154 launch_gui,
155 }))
156 }
157
158 pub fn default_socket_path() -> PathBuf {
160 if let Some(session) = crate::gui::live_sessions()
161 .into_iter()
162 .find(|s| s.alive && std::path::Path::new(&s.socket).exists())
163 {
164 return PathBuf::from(session.socket);
165 }
166 crate::app::paths::runtime_dir().join("approval.sock")
167 }
168
169 pub fn socket_path(&self) -> &std::path::Path {
170 &self.socket_path
171 }
172
173 pub fn answered_count(&self) -> u64 {
174 self.served.load(Ordering::Relaxed)
175 }
176
177 pub fn uptime(&self) -> Duration {
178 self.started.elapsed()
179 }
180
181 pub async fn ask(&self, request: ApprovalRequest) -> ConfirmOutcome {
185 self.ask_with_timeout(request, APPROVAL_IPC_TIMEOUT).await
186 }
187
188 async fn ask_with_timeout(
189 &self,
190 request: ApprovalRequest,
191 timeout: Duration,
192 ) -> ConfirmOutcome {
193 let id = request.id.clone();
194 let (tx, rx) = tokio::sync::oneshot::channel::<GrantChoice>();
195 {
196 let mut inner = self.inner.lock().unwrap();
197 if inner.slot.is_some() {
198 return ConfirmOutcome::Unavailable {
201 reason: "another approval request is already pending".into(),
202 };
203 }
204 inner.slot = Some(Slot {
205 request,
206 tx,
207 deadline: Instant::now() + timeout,
208 });
209 }
210 let _pending = PendingGuard {
211 inner: Arc::clone(&self.inner),
212 id: id.clone(),
213 };
214 let mut dialog = if self.launch_gui {
215 crate::gui::launch::open_prompt(&self.socket_path, &id)
216 } else {
217 None
218 };
219 let answer = async {
220 if let Some(child) = dialog.as_mut() {
221 tokio::select! {
222 biased;
223 result = rx => result.ok(),
224 _ = child.wait() => None,
225 }
226 } else {
227 rx.await.ok()
228 }
229 };
230 let outcome = match tokio::time::timeout(timeout, answer).await {
231 Ok(Some(choice)) => {
232 self.served.fetch_add(1, Ordering::Relaxed);
233 ConfirmOutcome::Chosen(choice)
234 }
235 Ok(None) => ConfirmOutcome::Unavailable {
236 reason: "approval window closed or companion dropped the request without a choice"
237 .into(),
238 },
239 Err(_elapsed) => {
240 ConfirmOutcome::Unavailable {
242 reason: format!("approval IPC timed out after {}s", timeout.as_secs()),
243 }
244 }
245 };
246 if let Some(child) = dialog.as_mut() {
249 let _ = tokio::time::timeout(Duration::from_millis(500), child.wait()).await;
250 }
251 outcome
252 }
253}
254
255impl Drop for ApprovalIpc {
256 fn drop(&mut self) {
257 self.inner.lock().unwrap().slot = None;
258 self.worker.abort();
259 let _ = std::fs::remove_file(&self.socket_path);
260 if let Some(session) = &self.session_file {
261 let _ = std::fs::remove_file(session);
262 }
263 }
264}
265
266fn peer_is_same_uid<Fd: std::os::unix::io::AsRawFd>(stream: &Fd) -> bool {
268 #[cfg(target_os = "macos")]
269 {
270 let mut cred: libc::xucred = unsafe { std::mem::zeroed() };
273 let mut len = std::mem::size_of_val(&cred) as libc::socklen_t;
274 let rc = unsafe {
275 libc::getsockopt(
276 stream.as_raw_fd(),
277 libc::SOL_LOCAL,
278 libc::LOCAL_PEERCRED,
279 &mut cred as *mut _ as *mut libc::c_void,
280 &mut len,
281 )
282 };
283 rc == 0 && cred.cr_uid == unsafe { libc::geteuid() }
284 }
285 #[cfg(target_os = "linux")]
286 {
287 let mut cred: libc::ucred = unsafe { std::mem::zeroed() };
288 let mut len = std::mem::size_of_val(&cred) as libc::socklen_t;
289 let rc = unsafe {
290 libc::getsockopt(
291 stream.as_raw_fd(),
292 libc::SOL_SOCKET,
293 libc::SO_PEERCRED,
294 &mut cred as *mut _ as *mut libc::c_void,
295 &mut len,
296 )
297 };
298 rc == 0 && cred.uid == unsafe { libc::geteuid() }
299 }
300 #[cfg(not(any(target_os = "macos", target_os = "linux")))]
301 {
302 let _ = stream;
304 false
305 }
306}
307
308async fn handle_connection(stream: UnixStream, inner: Arc<Mutex<Inner>>) {
309 if !peer_is_same_uid(&stream) {
311 return;
312 }
313 let (rd, mut wr) = stream.into_split();
314 let mut reader = BufReader::new(rd);
315 let Some(first) = read_line(&mut reader).await else {
316 return;
317 };
318 let Ok(cmd) = serde_json::from_str::<serde_json::Value>(&first) else {
319 return;
320 };
321 if cmd["op"].as_str() != Some("wait") {
322 return;
323 }
324
325 let deadline = Instant::now() + APPROVAL_IPC_TIMEOUT;
328 loop {
329 {
330 let guard = inner.lock().unwrap();
331 if guard.slot.is_some() || cmd["requestId"].is_string() || Instant::now() >= deadline {
332 break;
333 }
334 }
335 tokio::time::sleep(Duration::from_millis(25)).await;
336 }
337
338 let request = inner
339 .lock()
340 .unwrap()
341 .slot
342 .as_ref()
343 .filter(|slot| {
344 cmd["requestId"]
345 .as_str()
346 .is_none_or(|id| id == slot.request.id)
347 })
348 .map(|slot| serde_json::to_value(&slot.request).unwrap_or_default());
349 let Some(request) = request else {
350 let _ = write_line(&mut wr, &serde_json::json!({"op": "empty"})).await;
351 return;
352 };
353 if write_line(
354 &mut wr,
355 &serde_json::json!({
356 "op": "request",
357 "request": request,
358 }),
359 )
360 .await
361 .is_err()
362 {
363 return;
364 }
365
366 let request_id = request["id"].as_str().unwrap_or_default().to_string();
369 let expired = async {
370 loop {
371 let active = inner.lock().unwrap().slot.as_ref().is_some_and(|s| {
372 s.request.id == request_id && Instant::now() < s.deadline && !s.tx.is_closed()
373 });
374 if !active {
375 break;
376 }
377 tokio::time::sleep(Duration::from_millis(25)).await;
378 }
379 };
380 let line = tokio::select! {
381 biased;
382 _ = expired => {
383 let _ = write_line(&mut wr, &serde_json::json!({"op": "stale"})).await;
384 return;
385 }
386 line = read_line(&mut reader) => line,
387 };
388 let Some(line) = line else {
389 return;
390 };
391 let Ok(reply) = serde_json::from_str::<serde_json::Value>(&line) else {
392 return;
393 };
394 if reply["op"].as_str() != Some("reply") {
395 return;
396 }
397 enum Reply {
398 Choice(GrantChoice, Slot),
399 Stale,
400 BadChoice,
401 }
402 let decision = {
405 let mut guard = inner.lock().unwrap();
406 let id_matches = guard
407 .slot
408 .as_ref()
409 .map(|s| {
410 s.request.id == request_id
411 && s.request.id == reply["id"].as_str().unwrap_or("")
412 && Instant::now() < s.deadline
413 && !s.tx.is_closed()
414 })
415 .unwrap_or(false);
416 if !id_matches {
417 Reply::Stale
418 } else {
419 let choice = match reply["choice"].as_str() {
420 Some("once") => Some(GrantChoice::Once),
421 Some("session") => Some(GrantChoice::Session),
422 Some("decline") => Some(GrantChoice::Decline),
423 _ => None,
424 };
425 match choice {
426 Some(choice) => Reply::Choice(choice, guard.slot.take().unwrap()),
427 None => Reply::BadChoice,
428 }
429 }
430 };
431 match decision {
432 Reply::Stale => {
433 let _ = write_line(&mut wr, &serde_json::json!({"op": "stale"})).await;
434 }
435 Reply::BadChoice => {
436 let _ = write_line(&mut wr, &serde_json::json!({"op": "bad-choice"})).await;
437 }
438 Reply::Choice(choice, slot) => {
439 if write_line(&mut wr, &serde_json::json!({"op": "ok"}))
442 .await
443 .is_ok()
444 {
445 let _ = slot.tx.send(choice);
446 }
447 }
448 }
449}
450
451async fn read_line<R: tokio::io::AsyncBufRead + Unpin>(reader: &mut R) -> Option<String> {
452 let mut line = String::new();
453 reader.read_line(&mut line).await.ok()?;
454 let line = line.trim().to_string();
455 if line.is_empty() || line.len() > MAX_LINE {
456 return None;
457 }
458 Some(line)
459}
460
461async fn write_line<W: tokio::io::AsyncWrite + Unpin>(
462 stream: &mut W,
463 value: &serde_json::Value,
464) -> std::io::Result<()> {
465 stream
466 .write_all(serde_json::to_string(value).unwrap_or_default().as_bytes())
467 .await?;
468 stream.write_all(b"\n").await?;
469 stream.flush().await
470}
471
472fn iso_now() -> String {
473 time::OffsetDateTime::now_utc()
474 .format(&time::format_description::well_known::Rfc3339)
475 .unwrap_or_else(|_| "1970-01-01T00:00:00Z".into())
476}
477
478#[cfg(test)]
479mod tests {
480 use super::*;
481
482 fn request(id: &str) -> ApprovalRequest {
483 ApprovalRequest {
484 id: id.into(),
485 category: "write".into(),
486 connection: "local-dev".into(),
487 database: Some("app".into()),
488 tables: vec![],
489 snippet: "UPDATE items SET id = 2".into(),
490 }
491 }
492
493 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
494 async fn independent_server_sockets_do_not_replace_each_other() {
495 let dir = tempfile::TempDir::new().unwrap();
496 let first = ApprovalIpc::bind_at(dir.path().to_path_buf(), "p1.sock", false).unwrap();
497 let second = ApprovalIpc::bind_at(dir.path().to_path_buf(), "p2.sock", false).unwrap();
498 assert!(ApprovalIpc::bind_at(dir.path().to_path_buf(), "p1.sock", false).is_err());
499 let answer_first = tokio::spawn(approver_script(
500 first.socket_path().to_path_buf(),
501 "decline",
502 ));
503 let answer_second =
504 tokio::spawn(approver_script(second.socket_path().to_path_buf(), "once"));
505 let (a, b) = tokio::join!(first.ask(request("a")), second.ask(request("b")));
506 assert_eq!(a, ConfirmOutcome::Chosen(GrantChoice::Decline));
507 assert_eq!(b, ConfirmOutcome::Chosen(GrantChoice::Once));
508 answer_first.await.unwrap().unwrap();
509 answer_second.await.unwrap().unwrap();
510 drop(first);
511 assert!(second.socket_path().exists());
512 }
513
514 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
515 async fn timeout_notifies_idle_dialog_and_cancel_releases_slot() {
516 let dir = tempfile::TempDir::new().unwrap();
517 let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
518 let handle = crate::gui::companion::spawn_companion(ipc.socket_path().to_path_buf());
519 let result = ipc
520 .ask_with_timeout(request("expired"), Duration::from_millis(150))
521 .await;
522 assert!(matches!(result, ConfirmOutcome::Unavailable { .. }));
523 tokio::time::timeout(Duration::from_secs(2), async {
524 loop {
525 if handle
526 .events
527 .try_iter()
528 .any(|e| matches!(e, crate::gui::companion::CompanionEvent::Stale { .. }))
529 {
530 break;
531 }
532 tokio::time::sleep(Duration::from_millis(10)).await;
533 }
534 })
535 .await
536 .unwrap();
537 assert_eq!(ipc.answered_count(), 0);
538 let ask = tokio::spawn({
539 let ipc = Arc::clone(&ipc);
540 async move { ipc.ask(request("cancelled")).await }
541 });
542 tokio::time::sleep(Duration::from_millis(30)).await;
543 ask.abort();
544 let _ = ask.await;
545 assert!(ipc.inner.lock().unwrap().slot.is_none());
546 }
547
548 async fn approver_script(socket: PathBuf, choice: &'static str) -> Option<()> {
549 let stream = UnixStream::connect(socket).await.ok()?;
550 let (rd, mut wr) = stream.into_split();
551 let mut reader = BufReader::new(rd);
552 let hello = serde_json::json!({"op": "wait"});
553 write_line(&mut wr, &hello).await.ok()?;
554 let line = read_line(&mut reader).await?;
555 let msg: serde_json::Value = serde_json::from_str(&line).ok()?;
556 assert_eq!(msg["op"], "request", "{msg}");
557 let id = msg["request"]["id"].as_str()?.to_string();
558 let reply = serde_json::json!({"op": "reply", "id": id, "choice": choice});
559 write_line(&mut wr, &reply).await.ok()?;
560 let ack = read_line(&mut reader).await?;
561 let ack: serde_json::Value = serde_json::from_str(&ack).ok()?;
562 assert_eq!(ack["op"], "ok", "{ack}");
563 Some(())
564 }
565
566 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
567 async fn approve_once_round_trip() {
568 let dir = tempfile::TempDir::new().unwrap();
569 let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
570 let socket = ipc.socket_path().to_path_buf();
571 let handle = tokio::spawn(approver_script(socket, "once"));
572 let outcome = ipc
573 .ask(ApprovalRequest {
574 id: uuid::Uuid::new_v4().to_string(),
575 category: "write".into(),
576 connection: "c".into(),
577 database: Some("app".into()),
578 tables: vec!["app.users".into()],
579 snippet: "UPDATE users SET id = 2".into(),
580 })
581 .await;
582 assert_eq!(outcome, ConfirmOutcome::Chosen(GrantChoice::Once));
583 assert!(handle.await.unwrap().is_some());
584 assert_eq!(ipc.answered_count(), 1);
585 assert!(ipc.socket_path().exists());
586 drop(ipc);
587 assert!(!ipc_exists(dir.path()), "socket removed on drop");
588 }
589
590 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
591 async fn decline_and_session_choices() {
592 let dir = tempfile::TempDir::new().unwrap();
593 let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
594 for choice in ["decline", "session"] {
595 let socket = ipc.socket_path().to_path_buf();
596 let leaked: &'static str = Box::leak(choice.to_string().into_boxed_str());
597 let handle = tokio::spawn(approver_script(socket, leaked));
598 let outcome = ipc
599 .ask(ApprovalRequest {
600 id: uuid::Uuid::new_v4().to_string(),
601 category: "write".into(),
602 connection: "c".into(),
603 database: None,
604 tables: vec![],
605 snippet: "s".into(),
606 })
607 .await;
608 let expected = if choice == "decline" {
609 GrantChoice::Decline
610 } else {
611 GrantChoice::Session
612 };
613 assert_eq!(outcome, ConfirmOutcome::Chosen(expected));
614 assert!(handle.await.unwrap().is_some());
615 }
616 }
617
618 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
619 async fn no_companion_fails_closed_after_deadline() {
620 let dir = tempfile::TempDir::new().unwrap();
625 let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
626 let ipc2 = Arc::clone(&ipc);
627 let first = tokio::spawn(async move {
628 ipc2.ask(ApprovalRequest {
629 id: "first".into(),
630 category: "write".into(),
631 connection: "c".into(),
632 database: None,
633 tables: vec![],
634 snippet: "s".into(),
635 })
636 .await
637 });
638 tokio::time::sleep(Duration::from_millis(50)).await;
640 let second = ipc
641 .ask(ApprovalRequest {
642 id: "second".into(),
643 category: "write".into(),
644 connection: "c".into(),
645 database: None,
646 tables: vec![],
647 snippet: "s".into(),
648 })
649 .await;
650 assert!(
651 matches!(second, ConfirmOutcome::Unavailable { .. }),
652 "second concurrent ask fails closed: {second:?}"
653 );
654 first.abort();
655 }
656
657 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
658 async fn wrong_id_reply_is_stale() {
659 let dir = tempfile::TempDir::new().unwrap();
660 let ipc = ApprovalIpc::start_at(Some(dir.path().to_path_buf())).unwrap();
661 let socket = ipc.socket_path().to_path_buf();
662 let ask = tokio::spawn({
663 let ipc = Arc::clone(&ipc);
664 async move {
665 ipc.ask(ApprovalRequest {
666 id: "real-id".into(),
667 category: "write".into(),
668 connection: "c".into(),
669 database: None,
670 tables: vec![],
671 snippet: "s".into(),
672 })
673 .await
674 }
675 });
676 tokio::time::sleep(Duration::from_millis(100)).await;
677 let stream = UnixStream::connect(&socket).await.unwrap();
678 let (rd, mut wr) = stream.into_split();
679 let mut reader = BufReader::new(rd);
680 write_line(&mut wr, &serde_json::json!({"op": "wait"}))
681 .await
682 .unwrap();
683 let line = read_line(&mut reader).await.unwrap();
684 let msg: serde_json::Value = serde_json::from_str(&line).unwrap();
685 assert_eq!(msg["op"], "request");
686 write_line(
688 &mut wr,
689 &serde_json::json!({"op": "reply", "id": "forged", "choice": "once"}),
690 )
691 .await
692 .unwrap();
693 let ack = read_line(&mut reader).await.unwrap();
694 let ack: serde_json::Value = serde_json::from_str(&ack).unwrap();
695 assert_eq!(ack["op"], "stale", "{ack}");
696 ask.abort();
699 }
700
701 fn ipc_exists(dir: &std::path::Path) -> bool {
702 dir.join("approval.sock").exists()
703 }
704}