hx-remote 0.1.0

Open files in an existing Helix session through a tiny LSP bridge
Documentation
use crate::{
    SocketRequest, SocketResponse, absolute_path, read_lsp_message, show_document_params,
    write_lsp_message,
};
use serde_json::{Value, json};
use std::collections::{HashMap, VecDeque};
use std::fs;
use std::io::{self, BufReader, ErrorKind, Read, Write};
use std::os::unix::fs::{FileTypeExt, MetadataExt, PermissionsExt};
use std::os::unix::net::{UnixListener, UnixStream};
use std::path::{Path, PathBuf};
use std::sync::mpsc::{self, Receiver, Sender};
use std::thread;
use std::time::Duration;
use tempfile::{Builder, NamedTempFile};

const MAX_SOCKET_REQUEST_BYTES: u64 = 32 * 1024 * 1024;

enum Event {
    Lsp(Value),
    LspClosed,
    FatalInput(String),
    Open {
        request: SocketRequest,
        reply: Sender<Result<String, String>>,
    },
}

struct SocketGuard {
    path: PathBuf,
    device: u64,
    inode: u64,
}

impl Drop for SocketGuard {
    fn drop(&mut self) {
        let Ok(metadata) = fs::symlink_metadata(&self.path) else {
            return;
        };
        if metadata.dev() == self.device && metadata.ino() == self.inode {
            let _ = fs::remove_file(&self.path);
        }
    }
}

pub fn run_server(socket_path: PathBuf) -> Result<(), Box<dyn std::error::Error>> {
    let (listener, _socket_guard) = bind_socket(&socket_path)?;
    let (event_tx, event_rx) = mpsc::channel();

    spawn_lsp_reader(event_tx.clone());
    spawn_socket_listener(listener, event_tx);

    event_loop(event_rx)
}

fn bind_socket(socket_path: &Path) -> io::Result<(UnixListener, SocketGuard)> {
    if let Some(parent) = socket_path
        .parent()
        .filter(|parent| !parent.as_os_str().is_empty())
    {
        fs::create_dir_all(parent)?;
    }

    let listener = match UnixListener::bind(socket_path) {
        Ok(listener) => listener,
        Err(error) if error.kind() == ErrorKind::AddrInUse => {
            if UnixStream::connect(socket_path).is_ok() {
                return Err(io::Error::new(
                    ErrorKind::AddrInUse,
                    format!(
                        "another bridge is already listening on {}",
                        socket_path.display()
                    ),
                ));
            }

            let metadata = fs::symlink_metadata(socket_path)?;
            if !metadata.file_type().is_socket() {
                return Err(io::Error::new(
                    ErrorKind::AlreadyExists,
                    format!(
                        "refusing to replace non-socket path {}",
                        socket_path.display()
                    ),
                ));
            }
            fs::remove_file(socket_path)?;
            UnixListener::bind(socket_path)?
        }
        Err(error) => return Err(error),
    };

    fs::set_permissions(socket_path, fs::Permissions::from_mode(0o600))?;
    let metadata = fs::symlink_metadata(socket_path)?;
    let guard = SocketGuard {
        path: socket_path.to_path_buf(),
        device: metadata.dev(),
        inode: metadata.ino(),
    };
    Ok((listener, guard))
}

fn spawn_lsp_reader(event_tx: Sender<Event>) {
    thread::spawn(move || {
        let stdin = io::stdin();
        let mut reader = BufReader::new(stdin.lock());
        loop {
            match read_lsp_message(&mut reader) {
                Ok(Some(message)) => {
                    if event_tx.send(Event::Lsp(message)).is_err() {
                        return;
                    }
                }
                Ok(None) => {
                    let _ = event_tx.send(Event::LspClosed);
                    return;
                }
                Err(error) => {
                    let _ = event_tx.send(Event::FatalInput(error.to_string()));
                    return;
                }
            }
        }
    });
}

fn spawn_socket_listener(listener: UnixListener, event_tx: Sender<Event>) {
    thread::spawn(move || {
        for connection in listener.incoming() {
            match connection {
                Ok(stream) => {
                    let event_tx = event_tx.clone();
                    thread::spawn(move || handle_socket_connection(stream, event_tx));
                }
                Err(error) => {
                    eprintln!("hxr: socket accept failed: {error}");
                    return;
                }
            }
        }
    });
}

fn handle_socket_connection(mut stream: UnixStream, event_tx: Sender<Event>) {
    let result = (|| -> Result<String, String> {
        stream
            .set_read_timeout(Some(Duration::from_secs(10)))
            .map_err(|error| error.to_string())?;
        stream
            .set_write_timeout(Some(Duration::from_secs(10)))
            .map_err(|error| error.to_string())?;

        let mut request_bytes = Vec::new();
        BufReader::new(&stream)
            .take(MAX_SOCKET_REQUEST_BYTES + 1)
            .read_to_end(&mut request_bytes)
            .map_err(|error| error.to_string())?;
        if request_bytes.len() as u64 > MAX_SOCKET_REQUEST_BYTES {
            return Err(format!(
                "request exceeds the {} MiB limit",
                MAX_SOCKET_REQUEST_BYTES / 1024 / 1024
            ));
        }

        let request: SocketRequest = serde_json::from_slice(&request_bytes)
            .map_err(|error| format!("invalid bridge request: {error}"))?;
        let (reply_tx, reply_rx) = mpsc::channel();
        event_tx
            .send(Event::Open {
                request,
                reply: reply_tx,
            })
            .map_err(|_| "the LSP bridge has stopped".to_owned())?;
        reply_rx
            .recv_timeout(Duration::from_secs(10))
            .map_err(|_| "the LSP bridge did not accept the request in time".to_owned())?
    })();

    let response = match result {
        Ok(message) => SocketResponse::success(message),
        Err(message) => SocketResponse::error(message),
    };
    if serde_json::to_writer(&mut stream, &response).is_ok() {
        let _ = stream.write_all(b"\n");
        let _ = stream.flush();
    }
}

fn event_loop(event_rx: Receiver<Event>) -> Result<(), Box<dyn std::error::Error>> {
    let mut stdout = io::stdout().lock();
    let mut initialized = false;
    let mut shutting_down = false;
    let mut queued = VecDeque::new();
    let mut next_request_id = 1_u64;
    let mut pending_requests: HashMap<u64, String> = HashMap::new();
    let mut scratch_files: Vec<NamedTempFile> = Vec::new();

    while let Ok(event) = event_rx.recv() {
        match event {
            Event::Lsp(message) => {
                if let Some(method) = message.get("method").and_then(Value::as_str) {
                    match method {
                        "initialize" => {
                            if let Some(id) = message.get("id") {
                                let response = json!({
                                    "jsonrpc": "2.0",
                                    "id": id,
                                    "result": {
                                        "capabilities": {},
                                        "serverInfo": {
                                            "name": "hx-remote",
                                            "version": env!("CARGO_PKG_VERSION")
                                        }
                                    }
                                });
                                write_lsp_message(&mut stdout, &response)?;
                            }
                        }
                        "initialized" => {
                            initialized = true;
                            while let Some(request) = queued.pop_front() {
                                dispatch_open(
                                    request,
                                    &mut next_request_id,
                                    &mut stdout,
                                    &mut pending_requests,
                                    &mut scratch_files,
                                )?;
                            }
                        }
                        "shutdown" => {
                            shutting_down = true;
                            if let Some(id) = message.get("id") {
                                write_lsp_message(
                                    &mut stdout,
                                    &json!({"jsonrpc": "2.0", "id": id, "result": null}),
                                )?;
                            }
                        }
                        "exit" => return Ok(()),
                        _ => {}
                    }
                } else if let Some(id) = message.get("id").and_then(Value::as_u64)
                    && let Some(target) = pending_requests.remove(&id)
                {
                    let success = message
                        .get("result")
                        .and_then(|result| result.get("success"))
                        .and_then(Value::as_bool);
                    if success == Some(false) || message.get("error").is_some() {
                        eprintln!("hxr: Helix did not open {target}");
                    }
                }
            }
            Event::Open { request, reply } => {
                if shutting_down {
                    let _ = reply.send(Err("Helix is shutting down".into()));
                } else if initialized {
                    let result = dispatch_open(
                        request,
                        &mut next_request_id,
                        &mut stdout,
                        &mut pending_requests,
                        &mut scratch_files,
                    )
                    .map(|target| format!("sent {target} to Helix"))
                    .map_err(|error| error.to_string());
                    let _ = reply.send(result);
                } else {
                    queued.push_back(request);
                    let _ = reply.send(Ok("queued until Helix finishes initializing".into()));
                }
            }
            Event::LspClosed => return Ok(()),
            Event::FatalInput(error) => {
                return Err(format!("invalid LSP input: {error}").into());
            }
        }
    }
    Ok(())
}

fn dispatch_open(
    request: SocketRequest,
    next_request_id: &mut u64,
    stdout: &mut impl Write,
    pending_requests: &mut HashMap<u64, String>,
    scratch_files: &mut Vec<NamedTempFile>,
) -> io::Result<String> {
    let (path, line, column) = match request {
        SocketRequest::Open { path, line, column } => (absolute_path(&path)?, line, column),
        SocketRequest::OpenStdin { contents, name } => {
            let safe_name = sanitize_scratch_name(&name);
            let mut file = Builder::new()
                .prefix("hx-remote-")
                .suffix(&format!("-{safe_name}"))
                .tempfile()?;
            file.write_all(contents.as_bytes())?;
            file.flush()?;
            let path = file.path().to_path_buf();
            scratch_files.push(file);
            (path, None, None)
        }
    };

    let params = show_document_params(&path, line, column)
        .map_err(|message| io::Error::new(ErrorKind::InvalidInput, message))?;
    let id = *next_request_id;
    *next_request_id = next_request_id
        .checked_add(1)
        .ok_or_else(|| io::Error::other("LSP request id overflow"))?;
    let request = json!({
        "jsonrpc": "2.0",
        "id": id,
        "method": "window/showDocument",
        "params": params
    });
    write_lsp_message(stdout, &request)?;

    let target = path.display().to_string();
    pending_requests.insert(id, target.clone());
    Ok(target)
}

fn sanitize_scratch_name(name: &str) -> String {
    let name = Path::new(name)
        .file_name()
        .and_then(|name| name.to_str())
        .unwrap_or("stdin.txt");
    let sanitized: String = name
        .chars()
        .map(|character| {
            if character.is_ascii_alphanumeric() || matches!(character, '.' | '-' | '_') {
                character
            } else {
                '_'
            }
        })
        .take(80)
        .collect();
    if sanitized.is_empty() {
        "stdin.txt".into()
    } else {
        sanitized
    }
}

#[cfg(test)]
mod tests {
    use super::sanitize_scratch_name;

    #[test]
    fn scratch_names_cannot_escape_the_temp_directory() {
        assert_eq!(
            sanitize_scratch_name("../../my patch.diff"),
            "my_patch.diff"
        );
        assert_eq!(sanitize_scratch_name(""), "stdin.txt");
    }
}