use std::path::{Path, PathBuf};
use crossterm::event::{KeyCode, KeyEvent, KeyModifiers};
use ratatui::buffer::Buffer;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::UnixListener;
use tokio::sync::{mpsc, oneshot};
#[derive(Debug)]
pub enum ControlOp {
Screen(oneshot::Sender<String>),
Key(KeyEvent),
Command(String),
State(oneshot::Sender<String>),
Reload,
}
pub fn spawn_listener(path: PathBuf, tx: mpsc::UnboundedSender<ControlOp>) {
tokio::spawn(async move {
let _ = std::fs::remove_file(&path);
let listener = match UnixListener::bind(&path) {
Ok(l) => l,
Err(e) => {
tracing::error!(error = %e, path = %path.display(), "control socket bind failed");
return;
}
};
restrict_socket_perms(&path);
tracing::info!(path = %path.display(), "control socket listening");
let own_uid = socket_owner_uid(&path);
loop {
let (stream, _) = match listener.accept().await {
Ok(s) => s,
Err(e) => {
tracing::warn!(error = %e, "accept on control socket failed");
continue;
}
};
if !peer_is_owner(&stream, own_uid) {
tracing::warn!("control socket: rejected connection from another uid");
continue;
}
let tx2 = tx.clone();
tokio::spawn(async move {
let _ = handle_connection(stream, tx2).await;
});
}
});
}
#[cfg(unix)]
fn restrict_socket_perms(path: &Path) {
use std::os::unix::fs::PermissionsExt;
let _ = std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600));
}
#[cfg(not(unix))]
fn restrict_socket_perms(_path: &Path) {}
#[cfg(unix)]
fn socket_owner_uid(path: &Path) -> Option<u32> {
use std::os::unix::fs::MetadataExt;
std::fs::metadata(path).ok().map(|m| m.uid())
}
#[cfg(not(unix))]
fn socket_owner_uid(_path: &Path) -> Option<u32> {
None
}
pub(crate) fn uid_is_allowed(peer_uid: u32, own_uid: u32) -> bool {
peer_uid == own_uid || peer_uid == 0
}
#[cfg(unix)]
fn peer_is_owner(stream: &tokio::net::UnixStream, own_uid: Option<u32>) -> bool {
let Some(own) = own_uid else { return true };
match stream.peer_cred() {
Ok(cred) => uid_is_allowed(cred.uid(), own),
Err(_) => false,
}
}
#[cfg(not(unix))]
fn peer_is_owner(_stream: &tokio::net::UnixStream, _own_uid: Option<u32>) -> bool {
true
}
async fn handle_connection(
stream: tokio::net::UnixStream,
tx: mpsc::UnboundedSender<ControlOp>,
) -> std::io::Result<()> {
let (read_half, mut write_half) = stream.into_split();
let mut reader = BufReader::new(read_half);
let mut line = String::new();
reader.read_line(&mut line).await?;
let line = line.trim();
if line.is_empty() {
write_half.write_all(b"ERR empty request\n").await?;
return Ok(());
}
let (head, tail) = match line.split_once(' ') {
Some((h, t)) => (h, t),
None => (line, ""),
};
match head.to_ascii_uppercase().as_str() {
"SCREEN" => {
let (otx, orx) = oneshot::channel();
if tx.send(ControlOp::Screen(otx)).is_err() {
write_half.write_all(b"ERR app dropped channel\n").await?;
return Ok(());
}
match orx.await {
Ok(text) => {
write_half.write_all(text.as_bytes()).await?;
if !text.ends_with('\n') {
write_half.write_all(b"\n").await?;
}
}
Err(_) => {
write_half.write_all(b"ERR snapshot cancelled\n").await?;
}
}
}
"STATE" => {
let (otx, orx) = oneshot::channel();
if tx.send(ControlOp::State(otx)).is_err() {
write_half.write_all(b"ERR app dropped channel\n").await?;
return Ok(());
}
match orx.await {
Ok(text) => {
write_half.write_all(text.as_bytes()).await?;
write_half.write_all(b"\n").await?;
}
Err(_) => {
write_half.write_all(b"ERR state cancelled\n").await?;
}
}
}
"KEY" => match parse_key_spec(tail) {
Some(ke) => {
let _ = tx.send(ControlOp::Key(ke));
write_half.write_all(b"OK\n").await?;
}
None => {
write_half
.write_all(format!("ERR invalid key spec: {tail}\n").as_bytes())
.await?;
}
},
"RELOAD" => {
let _ = tx.send(ControlOp::Reload);
write_half.write_all(b"OK\n").await?;
}
"CMD" => {
let cmd = tail.trim().trim_start_matches(':').to_string();
if cmd.is_empty() {
write_half.write_all(b"ERR empty command\n").await?;
} else {
let _ = tx.send(ControlOp::Command(cmd));
write_half.write_all(b"OK\n").await?;
}
}
other => {
write_half
.write_all(
format!(
"ERR unknown op '{other}' (try: SCREEN | KEY <spec> | CMD <text> | STATE)\n"
)
.as_bytes(),
)
.await?;
}
}
Ok(())
}
pub(crate) fn render_buffer_as_text(buf: &Buffer) -> String {
let mut lines: Vec<String> = Vec::with_capacity(buf.area.height as usize);
for y in 0..buf.area.height {
let mut row = String::new();
for x in 0..buf.area.width {
let cell = &buf[(x, y)];
row.push_str(cell.symbol());
}
lines.push(row.trim_end().to_string());
}
lines.join("\n")
}
pub(crate) fn default_socket_path() -> PathBuf {
let mut p = crate::util::cache_dir();
p.push("control.sock");
p
}
pub(crate) fn parse_key_spec(spec: &str) -> Option<KeyEvent> {
let trimmed = spec.trim();
if trimmed.is_empty() {
return None;
}
let mut mods = KeyModifiers::NONE;
let mut code: Option<KeyCode> = None;
for piece in trimmed.split('+') {
let piece = piece.trim();
if piece.is_empty() {
continue;
}
let lower = piece.to_ascii_lowercase();
match lower.as_str() {
"ctrl" | "control" | "^" => mods |= KeyModifiers::CONTROL,
"shift" => mods |= KeyModifiers::SHIFT,
"alt" | "meta" | "option" => mods |= KeyModifiers::ALT,
"up" => code = Some(KeyCode::Up),
"down" => code = Some(KeyCode::Down),
"left" => code = Some(KeyCode::Left),
"right" => code = Some(KeyCode::Right),
"enter" | "return" => code = Some(KeyCode::Enter),
"esc" | "escape" => code = Some(KeyCode::Esc),
"tab" => code = Some(KeyCode::Tab),
"backtab" => code = Some(KeyCode::BackTab),
"backspace" => code = Some(KeyCode::Backspace),
"delete" | "del" => code = Some(KeyCode::Delete),
"home" => code = Some(KeyCode::Home),
"end" => code = Some(KeyCode::End),
"pageup" => code = Some(KeyCode::PageUp),
"pagedown" => code = Some(KeyCode::PageDown),
"space" => code = Some(KeyCode::Char(' ')),
_ => {
if let Some(num) = lower.strip_prefix('f').and_then(|n| n.parse::<u8>().ok()) {
if (1..=12).contains(&num) {
code = Some(KeyCode::F(num));
continue;
}
}
if let Some(inner) = piece
.strip_prefix("Char(")
.and_then(|s| s.strip_suffix(')'))
{
if let Some(c) = inner.chars().next() {
code = Some(KeyCode::Char(c));
continue;
}
}
if piece.chars().count() == 1 {
let c = piece.chars().next()?;
code = Some(KeyCode::Char(c));
}
}
}
}
if code == Some(KeyCode::Tab) && mods.contains(KeyModifiers::SHIFT) {
code = Some(KeyCode::BackTab);
}
code.map(|c| KeyEvent::new(c, mods))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_single_char_is_case_sensitive() {
let k = parse_key_spec("j").unwrap();
assert_eq!(k.code, KeyCode::Char('j'));
let k = parse_key_spec("J").unwrap();
assert_eq!(k.code, KeyCode::Char('J'));
}
#[test]
fn parse_arrow_keys() {
assert_eq!(parse_key_spec("Down").unwrap().code, KeyCode::Down);
assert_eq!(parse_key_spec("up").unwrap().code, KeyCode::Up);
}
#[test]
fn parse_ctrl_combinations() {
let k = parse_key_spec("Ctrl+R").unwrap();
assert_eq!(k.code, KeyCode::Char('R'));
assert!(k.modifiers.contains(KeyModifiers::CONTROL));
}
#[test]
fn parse_function_keys() {
assert_eq!(parse_key_spec("F2").unwrap().code, KeyCode::F(2));
assert_eq!(parse_key_spec("f12").unwrap().code, KeyCode::F(12));
assert!(parse_key_spec("F13").is_none());
}
#[test]
fn parse_explicit_char_form() {
let k = parse_key_spec("Char(:)").unwrap();
assert_eq!(k.code, KeyCode::Char(':'));
}
#[test]
fn parse_space_keyword() {
assert_eq!(parse_key_spec("Space").unwrap().code, KeyCode::Char(' '));
}
#[test]
fn parse_empty_is_none() {
assert!(parse_key_spec("").is_none());
assert!(parse_key_spec(" ").is_none());
}
}
#[cfg(test)]
mod peer_auth_tests {
use super::{peer_is_owner, uid_is_allowed};
#[test]
fn only_the_owner_and_root_may_drive_the_socket() {
assert!(uid_is_allowed(501, 501), "the owner");
assert!(uid_is_allowed(0, 501), "root could read the socket anyway");
assert!(uid_is_allowed(0, 0), "root owning it is still root");
assert!(
!uid_is_allowed(502, 501),
"another user must not drive this socket"
);
assert!(!uid_is_allowed(1, 501), "nor another system account");
assert!(!uid_is_allowed(65534, 501), "nor nobody(65534)");
}
#[tokio::test]
async fn peer_is_owner_reads_real_peer_credentials() {
let (a, _b) = tokio::net::UnixStream::pair().expect("socket pair");
let me = unsafe { libc::getuid() };
assert!(peer_is_owner(&a, Some(me)), "our own uid owns this socket");
assert!(
peer_is_owner(&a, None),
"no owner recorded (non-unix socket_owner_uid) means no check \
to make — the file permissions are the only gate there"
);
if me != 0 {
assert!(
!peer_is_owner(&a, Some(me.wrapping_add(1))),
"a socket owned by someone else must refuse us"
);
}
}
#[test]
fn the_listener_refuses_before_serving() {
let src = std::fs::read_to_string("src/control.rs").expect("read control.rs");
let listener = src
.split_once("\npub fn spawn_listener")
.expect("spawn_listener moved or was renamed")
.1;
let listener = listener.split("\n}\n").next().unwrap_or(listener);
assert!(
!listener.contains("mod peer_auth_tests"),
"the slice ran past the function into this test module, so it \
would be checking its own source"
);
assert!(
listener.contains("if !peer_is_owner(&stream, own_uid) {"),
"spawn_listener must refuse a connection whose peer is not the \
socket owner, BEFORE spawning handle_connection. Dropping the \
`!` inverts the gate and serves only other users."
);
assert!(
listener.contains("continue;"),
"and the refusal must skip the connection rather than fall \
through to serving it"
);
}
}
#[cfg(test)]
mod key_spec_tests {
use super::parse_key_spec;
use crossterm::event::{KeyCode, KeyModifiers};
#[test]
fn every_documented_key_name_parses() {
let docs = std::fs::read_to_string("docs/headless.md").expect("read headless.md");
let table = docs
.split_once("### `ctl key` spec vocabulary")
.expect("the ctl key vocabulary section is gone from the docs")
.1;
let table = table.split("\n##").next().unwrap_or(table);
let mut names: Vec<(String, bool)> = Vec::new();
for line in table.lines().filter(|l| l.trim_start().starts_with('|')) {
let is_modifier_row = line.contains("| modifiers |");
let mut rest = line;
while let Some((_, after)) = rest.split_once('`') {
let Some((tok, tail)) = after.split_once('`') else {
break;
};
rest = tail;
if tok.contains(' ') || tok.contains('…') || tok == "Char(x)" {
continue;
}
names.push((tok.to_string(), is_modifier_row));
}
}
assert!(
names.len() > 20,
"only {} names scraped from the docs table — the scrape is \
broken and this test would pass on nothing: {names:?}",
names.len()
);
for (name, is_modifier) in &names {
let spec = if *is_modifier {
format!("{name}+x")
} else {
name.clone()
};
assert!(
parse_key_spec(&spec).is_some(),
"docs/headless.md advertises `{name}` as a ctl key spec, \
and the parser rejects {spec:?}"
);
}
for (name, is_modifier) in &names {
if *is_modifier {
assert!(
parse_key_spec(name).is_none(),
"`{name}` is a modifier, not a key — it must not parse alone"
);
}
}
}
#[test]
fn every_key_name_the_parser_accepts_is_documented() {
let src = std::fs::read_to_string("src/control.rs").expect("read control.rs");
let body = src
.split_once("\npub(crate) fn parse_key_spec")
.expect("parse_key_spec moved or was renamed")
.1;
let body = body.split("\n}\n").next().unwrap_or(body);
assert!(
!body.contains("mod key_spec_tests"),
"the slice ran past the function into this test module"
);
let docs = std::fs::read_to_string("docs/headless.md").expect("read headless.md");
let mut undocumented = Vec::new();
for line in body.lines() {
let trimmed = line.trim();
if !trimmed.contains("=>") || !trimmed.starts_with('"') {
continue;
}
let pat = trimmed.split("=>").next().unwrap_or("");
for tok in pat.split('|') {
let name = tok.trim().trim_matches('"').trim();
if name.is_empty() {
continue;
}
if !docs.contains(&format!("`{name}`")) {
undocumented.push(name.to_string());
}
}
}
assert!(
undocumented.is_empty(),
"parse_key_spec accepts these and docs/headless.md doesn't \
mention them, so a script author can only find them by \
reading the source: {undocumented:?}"
);
}
#[test]
fn modifiers_accumulate_in_any_order() {
let k = parse_key_spec("ctrl+shift+alt+x").expect("parses");
assert_eq!(k.code, KeyCode::Char('x'));
assert!(k.modifiers.contains(KeyModifiers::CONTROL), "ctrl kept");
assert!(k.modifiers.contains(KeyModifiers::SHIFT), "shift kept");
assert!(k.modifiers.contains(KeyModifiers::ALT), "alt kept");
let k = parse_key_spec("alt+ctrl+r").expect("parses");
assert!(k.modifiers.contains(KeyModifiers::CONTROL));
assert!(k.modifiers.contains(KeyModifiers::ALT));
assert!(
!k.modifiers.contains(KeyModifiers::SHIFT),
"and only what was asked for"
);
}
#[test]
fn shift_tab_is_backtab() {
for spec in [
"shift+tab",
"Shift+Tab",
"tab+shift",
"SHIFT+TAB",
"backtab",
] {
let k = parse_key_spec(spec).unwrap_or_else(|| panic!("{spec} must parse"));
assert_eq!(
k.code,
KeyCode::BackTab,
"{spec} must be BackTab — the TUI moves BACKWARD on it and \
forward on Tab"
);
}
assert_eq!(parse_key_spec("tab").unwrap().code, KeyCode::Tab);
}
#[test]
fn a_spec_with_no_key_is_rejected() {
assert!(parse_key_spec("ctrl").is_none(), "modifiers alone");
assert!(parse_key_spec("ctrl+shift").is_none());
assert!(parse_key_spec("").is_none());
assert!(parse_key_spec(" ").is_none());
assert!(parse_key_spec("f13").is_none(), "out of the F1..F12 range");
assert!(parse_key_spec("f0").is_none());
assert!(parse_key_spec("pgup").is_none(), "not a name we accept");
}
}