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");
}
}