use std::io::Read;
use std::net::{IpAddr, SocketAddr};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use anyhow::{bail, Context, Result};
use serde_json::json;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::oneshot;
const DECISION_TIMEOUT: Duration = Duration::from_secs(180);
const MAX_REQUEST_BYTES: usize = 8 * 1024;
#[derive(Clone)]
pub struct ApprovalRequest {
pub tool: String,
pub risk: &'static str,
pub effect: Option<String>,
pub taint: Option<String>,
}
#[derive(Debug)]
struct Pending {
payload: String, tx: oneshot::Sender<bool>,
}
#[derive(Debug)]
pub struct ApprovalUi {
url: String,
token: String,
pending: Arc<Mutex<Option<Pending>>>,
opened: AtomicBool,
}
impl ApprovalUi {
pub async fn bind(addr: &str) -> Result<Arc<Self>> {
let sock: SocketAddr = addr
.parse()
.with_context(|| format!("`{addr}` is not a valid host:port"))?;
if !is_loopback(&sock.ip()) {
bail!(
"refusing to bind the approval UI to {sock}: anything that reaches this port can \
approve a mutation, so it is loopback-only by design"
);
}
let listener = TcpListener::bind(sock)
.await
.with_context(|| format!("could not bind {sock}"))?;
let bound = listener.local_addr()?;
let token = random_token()?;
let ui = Arc::new(ApprovalUi {
url: format!("http://{bound}/?t={token}"),
token,
pending: Arc::new(Mutex::new(None)),
opened: AtomicBool::new(false),
});
let serving = Arc::clone(&ui);
tokio::spawn(async move {
loop {
let Ok((stream, peer)) = listener.accept().await else {
continue;
};
if !is_loopback(&peer.ip()) {
continue;
}
let s = Arc::clone(&serving);
tokio::spawn(async move {
let _ = s.serve(stream).await;
});
}
});
Ok(ui)
}
pub fn url(&self) -> &str {
&self.url
}
pub async fn request(&self, req: ApprovalRequest) -> bool {
let (tx, rx) = oneshot::channel();
let payload = json!({
"tool": req.tool,
"risk": req.risk,
"effect": req.effect,
"taint": req.taint,
})
.to_string();
{
let mut slot = self.pending.lock().expect("pending lock");
*slot = Some(Pending { payload, tx });
}
if !self.opened.swap(true, Ordering::SeqCst) {
open_browser(&self.url);
}
eprintln!("⏸ waiting for approval at {}", self.url);
let decision = match tokio::time::timeout(DECISION_TIMEOUT, rx).await {
Ok(Ok(v)) => v,
_ => {
eprintln!(
"⏱ no decision within {}s — denied (the call stays a dry-run)",
DECISION_TIMEOUT.as_secs()
);
false
}
};
*self.pending.lock().expect("pending lock") = None;
decision
}
async fn serve(&self, mut stream: TcpStream) -> Result<()> {
let head = match read_head(&mut stream).await {
Some(h) => h,
None => return respond(&mut stream, 400, "text/plain", b"bad request").await,
};
let Some((method, target)) = request_line(&head) else {
return respond(&mut stream, 400, "text/plain", b"bad request").await;
};
if !host_is_loopback(&head) {
return respond(&mut stream, 403, "text/plain", b"bad host").await;
}
let (path, query) = target.split_once('?').unwrap_or((target, ""));
if !token_ok(query, &self.token) {
return respond(&mut stream, 404, "text/plain", b"not found").await;
}
match (method, path) {
("GET", "/") => {
respond(
&mut stream,
200,
"text/html; charset=utf-8",
PAGE.as_bytes(),
)
.await
}
("GET", "/pending") => {
let body = self
.pending
.lock()
.expect("pending lock")
.as_ref()
.map(|p| p.payload.clone())
.unwrap_or_else(|| "null".to_string());
respond(&mut stream, 200, "application/json", body.as_bytes()).await
}
("POST", "/decide") => {
if !header_contains(&head, "content-type", "application/json") {
return respond(&mut stream, 415, "text/plain", b"expected json").await;
}
let approve = read_body(&mut stream, &head)
.await
.and_then(|b| serde_json::from_slice::<serde_json::Value>(&b).ok())
.and_then(|v| v.get("approve").and_then(serde_json::Value::as_bool))
.unwrap_or(false);
let taken = self.pending.lock().expect("pending lock").take();
if let Some(p) = taken {
let _ = p.tx.send(approve);
}
eprintln!(
"{} decision from the approval UI",
if approve { "✅" } else { "🚫" }
);
respond(&mut stream, 200, "application/json", b"{\"ok\":true}").await
}
_ => respond(&mut stream, 404, "text/plain", b"not found").await,
}
}
}
fn is_loopback(ip: &IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => v4.is_loopback(),
IpAddr::V6(v6) => v6.is_loopback(),
}
}
fn random_token() -> Result<String> {
let mut buf = [0u8; 32];
std::fs::File::open("/dev/urandom")
.context("opening /dev/urandom for the approval token")?
.read_exact(&mut buf)
.context("reading the approval token")?;
Ok(buf.iter().map(|b| format!("{b:02x}")).collect())
}
fn constant_time_eq(a: &str, b: &str) -> bool {
let (a, b) = (a.as_bytes(), b.as_bytes());
if a.len() != b.len() {
return false;
}
a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
}
fn token_ok(query: &str, expected: &str) -> bool {
query
.split('&')
.filter_map(|kv| kv.split_once('='))
.any(|(k, v)| k == "t" && constant_time_eq(v, expected))
}
fn request_line(head: &str) -> Option<(&str, &str)> {
let first = head.lines().next()?;
let mut parts = first.split(' ');
Some((parts.next()?, parts.next()?))
}
fn header_value<'a>(head: &'a str, name: &str) -> Option<&'a str> {
head.lines()
.skip(1)
.filter_map(|l| l.split_once(':'))
.find(|(k, _)| k.trim().eq_ignore_ascii_case(name))
.map(|(_, v)| v.trim())
}
fn header_contains(head: &str, name: &str, needle: &str) -> bool {
header_value(head, name)
.map(|v| v.to_ascii_lowercase().contains(needle))
.unwrap_or(false)
}
fn host_is_loopback(head: &str) -> bool {
let Some(host) = header_value(head, "host") else {
return false; };
let name = host.rsplit_once(':').map(|(h, _)| h).unwrap_or(host);
let name = name.trim_start_matches('[').trim_end_matches(']');
name == "127.0.0.1" || name == "localhost" || name == "::1"
}
async fn read_head(stream: &mut TcpStream) -> Option<String> {
let mut buf = Vec::new();
let mut chunk = [0u8; 1024];
loop {
let n = stream.read(&mut chunk).await.ok()?;
if n == 0 {
return None;
}
buf.extend_from_slice(&chunk[..n]);
if buf.len() > MAX_REQUEST_BYTES {
return None;
}
if let Some(i) = find_head_end(&buf) {
let head = String::from_utf8(buf[..i].to_vec()).ok()?;
HEAD_REMAINDER.with(|r| *r.borrow_mut() = buf[i + 4..].to_vec());
return Some(head);
}
}
}
thread_local! {
static HEAD_REMAINDER: std::cell::RefCell<Vec<u8>> = const { std::cell::RefCell::new(Vec::new()) };
}
fn find_head_end(buf: &[u8]) -> Option<usize> {
buf.windows(4).position(|w| w == b"\r\n\r\n")
}
async fn read_body(stream: &mut TcpStream, head: &str) -> Option<Vec<u8>> {
let len: usize = header_value(head, "content-length")?.parse().ok()?;
if len > MAX_REQUEST_BYTES {
return None;
}
let mut body = HEAD_REMAINDER.with(|r| r.borrow_mut().split_off(0));
body.truncate(len);
while body.len() < len {
let mut chunk = vec![0u8; len - body.len()];
let n = stream.read(&mut chunk).await.ok()?;
if n == 0 {
return None;
}
body.extend_from_slice(&chunk[..n]);
}
Some(body)
}
async fn respond(stream: &mut TcpStream, status: u16, ctype: &str, body: &[u8]) -> Result<()> {
let reason = match status {
200 => "OK",
400 => "Bad Request",
403 => "Forbidden",
404 => "Not Found",
415 => "Unsupported Media Type",
_ => "Error",
};
let head = format!(
"HTTP/1.1 {status} {reason}\r\nContent-Type: {ctype}\r\nContent-Length: {}\r\n\
Cache-Control: no-store\r\nX-Content-Type-Options: nosniff\r\nConnection: close\r\n\r\n",
body.len()
);
stream.write_all(head.as_bytes()).await?;
stream.write_all(body).await?;
stream.flush().await?;
Ok(())
}
fn open_browser(url: &str) {
if std::env::var_os("FOREGUARD_NO_OPEN").is_some() {
return;
}
let opener = if cfg!(target_os = "macos") {
"open"
} else {
"xdg-open"
};
let _ = std::process::Command::new(opener)
.arg(url)
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.spawn();
}
const PAGE: &str = include_str!("ui.html");
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_non_loopback_bind_is_refused() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
let err = rt.block_on(ApprovalUi::bind("0.0.0.0:0")).unwrap_err();
assert!(
err.to_string().contains("loopback-only"),
"expected a loopback refusal, got: {err}"
);
}
#[test]
fn tokens_are_unique_and_long() {
let a = random_token().unwrap();
let b = random_token().unwrap();
assert_eq!(a.len(), 64, "32 bytes hex encoded");
assert_ne!(a, b, "two tokens in a row must not match");
}
#[test]
fn token_comparison_rejects_prefixes_and_wrong_lengths() {
assert!(constant_time_eq("abc", "abc"));
assert!(!constant_time_eq("abc", "abd"));
assert!(!constant_time_eq("abc", "ab"));
assert!(!constant_time_eq("ab", "abc"));
assert!(!constant_time_eq("", "a"));
}
#[test]
fn only_an_exact_token_in_the_query_is_accepted() {
assert!(token_ok("t=secret", "secret"));
assert!(token_ok("x=1&t=secret", "secret"));
assert!(!token_ok("t=secre", "secret"));
assert!(!token_ok("t=secrets", "secret"));
assert!(!token_ok("token=secret", "secret"));
assert!(!token_ok("", "secret"));
}
#[test]
fn host_must_be_a_loopback_literal() {
let h = |v: &str| format!("GET / HTTP/1.1\r\nHost: {v}");
assert!(host_is_loopback(&h("127.0.0.1:7878")));
assert!(host_is_loopback(&h("localhost:7878")));
assert!(host_is_loopback(&h("[::1]:7878")));
assert!(!host_is_loopback(&h("evil.example.com:7878")));
assert!(!host_is_loopback(&h("127.0.0.1.evil.com")));
assert!(!host_is_loopback("GET / HTTP/1.1\r\nAccept: */*"));
}
async fn http(addr: &str, raw: &str) -> (u16, String, String) {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut s = tokio::net::TcpStream::connect(addr).await.expect("connect");
s.write_all(raw.as_bytes()).await.expect("write");
let mut buf = Vec::new();
s.read_to_end(&mut buf).await.expect("read");
let text = String::from_utf8_lossy(&buf).into_owned();
let status = text
.split(' ')
.nth(1)
.and_then(|c| c.parse().ok())
.unwrap_or(0);
let (head, body) = text.split_once("\r\n\r\n").unwrap_or((&text, ""));
(status, head.to_string(), body.to_string())
}
async fn ui_on_free_port() -> (std::sync::Arc<ApprovalUi>, String, String) {
let ui = ApprovalUi::bind("127.0.0.1:0").await.expect("bind");
let rest = ui.url().trim_start_matches("http://").to_string();
let (addr, token) = rest.split_once("/?t=").expect("url shape");
(
std::sync::Arc::clone(&ui),
addr.to_string(),
token.to_string(),
)
}
fn get(path: &str, host: &str) -> String {
format!("GET {path} HTTP/1.1\r\nHost: {host}\r\nConnection: close\r\n\r\n")
}
#[tokio::test]
async fn without_the_token_nothing_is_reachable() {
let (_ui, addr, token) = ui_on_free_port().await;
for path in ["/", "/pending", "/decide"] {
let (status, _, _) = http(&addr, &get(path, &addr)).await;
assert_eq!(status, 404, "{path} answered without a token");
}
let (status, _, body) = http(&addr, &get(&format!("/?t={token}"), &addr)).await;
assert_eq!(status, 200);
assert!(body.contains("approval required") || body.contains("foreguard"));
}
#[tokio::test]
async fn a_non_loopback_host_header_is_refused_even_with_the_token() {
let (_ui, addr, token) = ui_on_free_port().await;
let raw = get(&format!("/pending?t={token}"), "evil.example.com");
let (status, _, _) = http(&addr, &raw).await;
assert_eq!(status, 403, "a rebound host was served");
}
#[tokio::test]
async fn no_cors_header_is_ever_sent() {
let (_ui, addr, token) = ui_on_free_port().await;
let (_, head, _) = http(&addr, &get(&format!("/?t={token}"), &addr)).await;
assert!(
!head
.to_ascii_lowercase()
.contains("access-control-allow-origin"),
"a CORS header would let a cross-origin page read approvals:\n{head}"
);
}
#[tokio::test]
async fn pending_is_null_until_something_needs_approving() {
let (_ui, addr, token) = ui_on_free_port().await;
let (status, _, body) = http(&addr, &get(&format!("/pending?t={token}"), &addr)).await;
assert_eq!(status, 200);
assert_eq!(body.trim(), "null");
}
#[tokio::test]
async fn an_approval_round_trips_and_the_decision_reaches_the_caller() {
let (ui, addr, token) = ui_on_free_port().await;
let asking = {
let ui = std::sync::Arc::clone(&ui);
tokio::spawn(async move {
ui.request(ApprovalRequest {
tool: "delete_file".into(),
risk: "high",
effect: Some("deletes /etc/passwd".into()),
taint: Some("attacker@evil.com".into()),
})
.await
})
};
let mut body = String::new();
for _ in 0..100 {
let (_, _, b) = http(&addr, &get(&format!("/pending?t={token}"), &addr)).await;
if b.trim() != "null" {
body = b;
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert!(
body.contains("delete_file"),
"pending never appeared: {body}"
);
assert!(
body.contains("deletes /etc/passwd"),
"effect missing: {body}"
);
assert!(body.contains("attacker@evil.com"), "taint missing: {body}");
let payload = r#"{"approve":true}"#;
let raw = format!(
"POST /decide?t={token} HTTP/1.1\r\nHost: {addr}\r\nContent-Type: application/json\r\n\
Content-Length: {}\r\nConnection: close\r\n\r\n{payload}",
payload.len()
);
let (status, _, _) = http(&addr, &raw).await;
assert_eq!(status, 200);
assert!(
asking.await.expect("join"),
"the approval did not reach the caller"
);
}
#[tokio::test]
async fn a_denial_reaches_the_caller_as_false() {
let (ui, addr, token) = ui_on_free_port().await;
let asking = {
let ui = std::sync::Arc::clone(&ui);
tokio::spawn(async move {
ui.request(ApprovalRequest {
tool: "write_file".into(),
risk: "medium",
effect: None,
taint: None,
})
.await
})
};
for _ in 0..100 {
let (_, _, b) = http(&addr, &get(&format!("/pending?t={token}"), &addr)).await;
if b.trim() != "null" {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
let payload = r#"{"approve":false}"#;
let raw = format!(
"POST /decide?t={token} HTTP/1.1\r\nHost: {addr}\r\nContent-Type: application/json\r\n\
Content-Length: {}\r\nConnection: close\r\n\r\n{payload}",
payload.len()
);
http(&addr, &raw).await;
assert!(
!asking.await.expect("join"),
"a denial came back as approval"
);
}
#[tokio::test]
async fn a_form_content_type_cannot_approve() {
let (_ui, addr, token) = ui_on_free_port().await;
let payload = "approve=true";
let raw = format!(
"POST /decide?t={token} HTTP/1.1\r\nHost: {addr}\r\n\
Content-Type: application/x-www-form-urlencoded\r\n\
Content-Length: {}\r\nConnection: close\r\n\r\n{payload}",
payload.len()
);
let (status, _, _) = http(&addr, &raw).await;
assert_eq!(status, 415, "a simple cross-origin POST was accepted");
}
#[tokio::test]
async fn an_unparseable_body_denies() {
let (ui, addr, token) = ui_on_free_port().await;
let asking = {
let ui = std::sync::Arc::clone(&ui);
tokio::spawn(async move {
ui.request(ApprovalRequest {
tool: "rm".into(),
risk: "high",
effect: None,
taint: None,
})
.await
})
};
for _ in 0..100 {
let (_, _, b) = http(&addr, &get(&format!("/pending?t={token}"), &addr)).await;
if b.trim() != "null" {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
let payload = "not json at all";
let raw = format!(
"POST /decide?t={token} HTTP/1.1\r\nHost: {addr}\r\nContent-Type: application/json\r\n\
Content-Length: {}\r\nConnection: close\r\n\r\n{payload}",
payload.len()
);
http(&addr, &raw).await;
assert!(!asking.await.expect("join"), "garbage was read as approval");
}
#[test]
fn headers_are_matched_case_insensitively() {
let head = "POST /decide HTTP/1.1\r\nCONTENT-TYPE: Application/JSON\r\nHost: localhost";
assert!(header_contains(head, "content-type", "application/json"));
assert!(!header_contains(head, "content-type", "text/plain"));
}
}