use crate::{
agent::cancellation::AgentCancellation,
config::LspServerConfig,
mcp::jsonrpc::{
ErrorData, JsonRpcError, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest,
METHOD_NOT_FOUND, RequestId,
},
tools::process::{
PIPE_READER_JOIN_TIMEOUT, PipeReaderHandle, recv_pipe_reader_with_timeout,
spawn_bounded_pipe_reader, terminate_child_tree_and_wait,
},
};
use anyhow::{Context, bail};
use serde_json::{Value, json};
use std::{
collections::HashMap,
env,
io::{BufRead, BufReader, BufWriter, Write},
path::Path,
process::{Child, ChildStdin, Command, Stdio},
sync::{
Arc, Mutex,
atomic::{AtomicBool, AtomicI64, Ordering},
mpsc::{self, Receiver, SyncSender},
},
thread::{self, JoinHandle},
time::{Duration, Instant},
};
const MAX_LSP_HEADER_BYTES: usize = 8192;
const MAX_LSP_BODY_BYTES: usize = 1_048_576;
const MAX_LSP_STDERR_BYTES: usize = 8192;
type PendingMap = Arc<Mutex<HashMap<RequestId, SyncSender<anyhow::Result<Value>>>>>;
pub(crate) struct LspStdioTransport {
writer: Arc<Mutex<BufWriter<ChildStdin>>>,
child: Arc<Mutex<Child>>,
reader: Option<JoinHandle<()>>,
reader_done: Option<Receiver<()>>,
stderr_receiver: Option<PipeReaderHandle>,
pending: PendingMap,
notifications: Receiver<JsonRpcNotification>,
next_id: AtomicI64,
closed: Arc<AtomicBool>,
}
#[derive(Clone)]
struct LspCleanupHandle {
child: Arc<Mutex<Child>>,
}
struct PendingRequestCleanup {
pending: PendingMap,
id: RequestId,
}
impl Drop for PendingRequestCleanup {
fn drop(&mut self) {
if let Ok(mut pending) = self.pending.lock() {
pending.remove(&self.id);
}
}
}
impl LspCleanupHandle {
fn terminate(&self) {
if let Ok(mut child) = self.child.lock() {
let _ = terminate_child_tree_and_wait(&mut child);
}
}
}
impl LspStdioTransport {
pub(crate) fn spawn(config: &LspServerConfig, cwd: &Path) -> anyhow::Result<Self> {
if config.command.trim().is_empty() {
bail!("LSP stdio command must not be empty");
}
let mut command = Command::new(&config.command);
command
.args(&config.args)
.current_dir(cwd)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.env_clear();
if let Some(path) = env::var_os("PATH") {
command.env("PATH", path);
}
#[cfg(unix)]
{
use std::os::unix::process::CommandExt;
command.process_group(0);
}
let mut child = command.spawn().context("spawning LSP server")?;
let stdin = child.stdin.take().context("LSP stdin pipe unavailable")?;
let stdout = child.stdout.take().context("LSP stdout pipe unavailable")?;
let stderr = child.stderr.take().context("LSP stderr pipe unavailable")?;
let child = Arc::new(Mutex::new(child));
let writer = Arc::new(Mutex::new(BufWriter::new(stdin)));
let pending: PendingMap = Arc::new(Mutex::new(HashMap::new()));
let (notification_sender, notifications) = mpsc::sync_channel(256);
let (reader_done_sender, reader_done) = mpsc::sync_channel(1);
let closed = Arc::new(AtomicBool::new(false));
let cleanup = LspCleanupHandle {
child: Arc::clone(&child),
};
let reader_pending = Arc::clone(&pending);
let reader_writer = Arc::clone(&writer);
let reader_closed = Arc::clone(&closed);
let reader = thread::spawn(move || {
let mut reader = BufReader::new(stdout);
loop {
let value = match read_lsp_message(&mut reader) {
Ok(value) => value,
Err(error) => {
reader_closed.store(true, Ordering::SeqCst);
fail_pending(
&reader_pending,
anyhow::anyhow!("LSP reader failed: {error}"),
);
cleanup.terminate();
break;
}
};
if value.get("method").is_some() && value.get("id").is_some() {
match serde_json::from_value::<JsonRpcRequest>(value) {
Ok(request) => {
let _ = respond_to_server_request(&reader_writer, request);
continue;
}
Err(error) => {
reader_closed.store(true, Ordering::SeqCst);
fail_pending(
&reader_pending,
anyhow::anyhow!("LSP malformed JSON-RPC request: {error}"),
);
cleanup.terminate();
break;
}
}
}
let message = match serde_json::from_value::<JsonRpcMessage>(value.clone()) {
Ok(message) => message,
Err(error) => {
reader_closed.store(true, Ordering::SeqCst);
fail_pending(
&reader_pending,
anyhow::anyhow!("LSP malformed JSON-RPC: {error}"),
);
cleanup.terminate();
break;
}
};
match message {
JsonRpcMessage::Response(response) => {
if let Ok(mut pending) = reader_pending.lock()
&& let Some(sender) = pending.remove(&response.id)
{
let _ = sender.send(Ok(response.result));
}
}
JsonRpcMessage::Error(error) => {
if let Some(id) = error.id
&& let Ok(mut pending) = reader_pending.lock()
&& let Some(sender) = pending.remove(&id)
{
let ErrorData { code, message, .. } = error.error;
let _ = sender
.send(Err(anyhow::anyhow!("LSP protocol error {code}: {message}")));
}
}
JsonRpcMessage::Notification(notification) => {
let _ = notification_sender.try_send(notification);
}
JsonRpcMessage::Request(request) => {
let _ = respond_to_server_request(&reader_writer, request);
}
}
}
let _ = reader_done_sender.send(());
});
let stderr_truncated = Arc::new(AtomicBool::new(false));
let stderr_receiver = Some(spawn_bounded_pipe_reader(
stderr,
MAX_LSP_STDERR_BYTES,
Arc::clone(&stderr_truncated),
));
Ok(Self {
writer,
child,
reader: Some(reader),
reader_done: Some(reader_done),
stderr_receiver,
pending,
notifications,
next_id: AtomicI64::new(1),
closed,
})
}
pub(crate) fn request(
&self,
method: &str,
params: Value,
timeout: Duration,
cancellation: &AgentCancellation,
) -> anyhow::Result<Value> {
if self.closed.load(Ordering::SeqCst) {
bail!("LSP server is closed");
}
let id = RequestId::Number(self.next_id.fetch_add(1, Ordering::SeqCst));
let (sender, receiver) = mpsc::sync_channel(1);
self.pending
.lock()
.map_err(|_| anyhow::anyhow!("LSP pending request map poisoned"))?
.insert(id.clone(), sender);
let _pending_cleanup = PendingRequestCleanup {
pending: Arc::clone(&self.pending),
id: id.clone(),
};
let message = serde_json::to_value(JsonRpcRequest::new(id.clone(), method, Some(params)))?;
self.write_value(&message)?;
let started = Instant::now();
loop {
cancellation.check()?;
let remaining = timeout.saturating_sub(started.elapsed());
if remaining.is_zero() {
bail!("LSP request timed out after {}ms", timeout.as_millis());
}
match receiver.recv_timeout(remaining.min(Duration::from_millis(50))) {
Ok(result) => return result,
Err(mpsc::RecvTimeoutError::Timeout) => continue,
Err(mpsc::RecvTimeoutError::Disconnected) => {
bail!("LSP reader disconnected before response");
}
}
}
}
pub(crate) fn notify(&self, method: &str, params: Value) -> anyhow::Result<()> {
let message = serde_json::to_value(JsonRpcNotification::new(method, Some(params)))?;
self.write_value(&message)
}
pub(crate) fn try_recv_notification(&self) -> Option<JsonRpcNotification> {
self.notifications.try_recv().ok()
}
pub(crate) fn is_closed(&self) -> bool {
self.closed.load(Ordering::SeqCst)
}
pub(crate) fn shutdown(&mut self, timeout: Duration) {
let cancellation = AgentCancellation::default();
let _ = self.request("shutdown", Value::Null, timeout, &cancellation);
let _ = self.notify("exit", Value::Null);
self.terminate();
if let Some(handle) = self.reader.take()
&& let Some(done) = self.reader_done.take()
{
match done.recv_timeout(PIPE_READER_JOIN_TIMEOUT) {
Ok(()) | Err(mpsc::RecvTimeoutError::Disconnected) => {
let _ = handle.join();
}
Err(mpsc::RecvTimeoutError::Timeout) => {}
}
}
if let Some(receiver) = self.stderr_receiver.take() {
let _ = recv_pipe_reader_with_timeout(receiver);
}
}
fn write_value(&self, value: &Value) -> anyhow::Result<()> {
let mut writer = self
.writer
.lock()
.map_err(|_| anyhow::anyhow!("LSP writer lock poisoned"))?;
write_lsp_message(&mut *writer, value)
}
fn terminate(&self) {
if let Ok(mut child) = self.child.lock() {
let _ = terminate_child_tree_and_wait(&mut child);
}
}
#[cfg(test)]
pub(crate) fn is_process_running(&self) -> bool {
self.child
.lock()
.ok()
.and_then(|mut child| child.try_wait().ok().flatten())
.is_none()
}
}
impl Drop for LspStdioTransport {
fn drop(&mut self) {
self.shutdown(Duration::from_millis(100));
}
}
fn respond_to_server_request(
writer: &Arc<Mutex<BufWriter<ChildStdin>>>,
request: JsonRpcRequest,
) -> anyhow::Result<()> {
let value = if matches!(
request.method.as_str(),
"workspace/configuration" | "window/workDoneProgress/create"
) {
json!({"jsonrpc":"2.0", "id": request.id, "result": null})
} else {
serde_json::to_value(JsonRpcError {
jsonrpc: "2.0".to_string(),
id: Some(request.id),
error: ErrorData {
code: METHOD_NOT_FOUND,
message: "Method not found".to_string(),
data: None,
},
})?
};
let mut writer = writer
.lock()
.map_err(|_| anyhow::anyhow!("LSP writer lock poisoned"))?;
write_lsp_message(&mut *writer, &value)
}
fn fail_pending(pending: &PendingMap, error: anyhow::Error) {
if let Ok(mut pending) = pending.lock() {
let message = error.to_string();
for (_, sender) in pending.drain() {
let _ = sender.send(Err(anyhow::anyhow!(message.clone())));
}
}
}
pub(crate) fn write_lsp_message(writer: &mut dyn Write, value: &Value) -> anyhow::Result<()> {
let body = serde_json::to_vec(value)?;
write!(writer, "Content-Length: {}\r\n\r\n", body.len())?;
writer.write_all(&body)?;
writer.flush()?;
Ok(())
}
pub(crate) fn read_lsp_message(reader: &mut dyn BufRead) -> anyhow::Result<Value> {
let mut header = Vec::new();
loop {
let mut line = Vec::new();
let count = reader.read_until(b'\n', &mut line)?;
if count == 0 {
bail!("LSP stdout closed");
}
header.extend_from_slice(&line);
if header.len() > MAX_LSP_HEADER_BYTES {
bail!("LSP header exceeded {MAX_LSP_HEADER_BYTES} bytes");
}
if header.ends_with(b"\r\n\r\n") || header.ends_with(b"\n\n") {
break;
}
}
let header = String::from_utf8(header)?;
let content_length = parse_content_length_header(&header)?;
if content_length > MAX_LSP_BODY_BYTES {
bail!("LSP body exceeded {MAX_LSP_BODY_BYTES} bytes");
}
let mut body = vec![0; content_length];
reader.read_exact(&mut body)?;
Ok(serde_json::from_slice(&body)?)
}
pub(crate) fn parse_content_length_header(headers: &str) -> anyhow::Result<usize> {
for line in headers.lines() {
let Some((name, value)) = line.split_once(':') else {
continue;
};
if name.eq_ignore_ascii_case("Content-Length") {
let length = value.trim().parse::<usize>()?;
return Ok(length);
}
}
bail!("missing Content-Length header")
}
#[cfg(test)]
mod tests {
use super::*;
use std::{io::Cursor, thread};
fn lsp_config(command: &str, args: Vec<String>) -> LspServerConfig {
LspServerConfig {
command: command.to_string(),
args,
enabled: true,
}
}
#[test]
fn parse_content_length_header_extracts_byte_count() {
assert_eq!(
parse_content_length_header("Content-Length: 12\r\n\r\n").unwrap(),
12
);
assert_eq!(
parse_content_length_header("X: y\r\ncontent-length: 7\r\n\r\n").unwrap(),
7
);
}
#[test]
fn parse_content_length_header_rejects_malformed_headers() {
assert!(parse_content_length_header("X: y\r\n\r\n").is_err());
assert!(parse_content_length_header("Content-Length: nope\r\n\r\n").is_err());
}
#[test]
fn write_lsp_message_writes_exact_content_length_frame() {
let mut bytes = Vec::new();
write_lsp_message(&mut bytes, &json!({"jsonrpc":"2.0","method":"x"})).unwrap();
let text = String::from_utf8(bytes).unwrap();
assert!(text.starts_with("Content-Length: 30\r\n\r\n"), "{text}");
assert!(
text.ends_with(r#"{"jsonrpc":"2.0","method":"x"}"#),
"{text}"
);
}
#[test]
fn read_lsp_message_reads_exact_body() {
let mut reader = Cursor::new(b"Content-Length: 11\r\n\r\n{\"ok\":true}trailing");
assert_eq!(read_lsp_message(&mut reader).unwrap(), json!({"ok": true}));
}
#[test]
fn read_lsp_message_rejects_oversized_body() {
let mut reader = Cursor::new(format!(
"Content-Length: {}\r\n\r\n",
MAX_LSP_BODY_BYTES + 1
));
let error = read_lsp_message(&mut reader).unwrap_err().to_string();
assert!(error.contains("exceeded"), "{error}");
}
#[test]
fn spawn_filters_environment_to_path_only() {
let script = "import os, time; print('Content-Length: '+str(len('{\\\"jsonrpc\\\":\\\"2.0\\\",\\\"method\\\":\\\"env\\\",\\\"params\\\":'+__import__('json').dumps(dict(os.environ))+'}'))+'\\r\\n\\r\\n'+'{\\\"jsonrpc\\\":\\\"2.0\\\",\\\"method\\\":\\\"env\\\",\\\"params\\\":'+__import__('json').dumps(dict(os.environ))+'}', end='', flush=True); time.sleep(5)";
let transport = LspStdioTransport::spawn(
&lsp_config(
"python3",
vec!["-u".to_string(), "-c".to_string(), script.to_string()],
),
Path::new("."),
)
.unwrap();
let notification = wait_for_notification(&transport, Duration::from_secs(2)).unwrap();
let env = notification.params.unwrap();
assert!(env.get("PATH").is_some());
assert!(env.get("OPENAI_API_KEY").is_none());
assert!(env.get("MC_API_KEY").is_none());
assert!(env.get("ANTHROPIC_API_KEY").is_none());
}
#[test]
fn fake_lsp_helper_reports_only_allowlisted_environment() {
let transport = LspStdioTransport::spawn(
&lsp_config(
"python3",
vec![
"-u".to_string(),
"tests/support/fake_lsp_server.py".to_string(),
"env_report".to_string(),
],
),
Path::new("."),
)
.unwrap();
let _ = transport
.request(
"initialize",
json!({"rootUri":"file:///tmp"}),
Duration::from_secs(1),
&AgentCancellation::default(),
)
.unwrap();
transport.notify("initialized", json!({})).unwrap();
let notification = wait_for_notification(&transport, Duration::from_secs(2)).unwrap();
let env = notification.params.unwrap().as_object().unwrap().clone();
assert!(env.contains_key("PATH"));
for secret_name in ["OPENAI_API_KEY", "MC_API_KEY", "ANTHROPIC_API_KEY"] {
assert!(
!env.contains_key(secret_name),
"leaked env secret key: {secret_name}"
);
}
assert!(env.keys().all(|key| {
let key = key.to_ascii_lowercase();
!key.contains("token") && !key.contains("bearer") && !key.contains("key")
}));
}
#[test]
fn drop_kills_child_process() {
let transport = LspStdioTransport::spawn(
&lsp_config("sh", vec!["-c".to_string(), "sleep 30".to_string()]),
Path::new("."),
)
.unwrap();
assert!(transport.is_process_running());
let child = Arc::clone(&transport.child);
drop(transport);
let started = Instant::now();
while started.elapsed() < Duration::from_secs(2) {
if child
.lock()
.ok()
.and_then(|mut child| child.try_wait().ok().flatten())
.is_some()
{
return;
}
thread::sleep(Duration::from_millis(20));
}
panic!("child still running after drop");
}
#[test]
fn canceled_request_removes_pending_entry_and_server_stays_usable() {
let transport = LspStdioTransport::spawn(
&lsp_config(
"python3",
vec![
"-u".to_string(),
"tests/support/fake_lsp_server.py".to_string(),
"slow_references".to_string(),
],
),
Path::new("."),
)
.unwrap();
let transport = Mutex::new(transport);
transport
.lock()
.unwrap()
.request(
"initialize",
json!({"rootUri":"file:///tmp"}),
Duration::from_secs(1),
&AgentCancellation::default(),
)
.unwrap();
let (cancellation, handle) = AgentCancellation::default().child_token();
let pending = Arc::clone(&transport.lock().unwrap().pending);
thread::scope(|scope| {
let request = scope.spawn(|| {
transport.lock().unwrap().request(
"textDocument/references",
json!({"textDocument":{"uri":"file:///tmp/file.rs"}}),
Duration::from_secs(2),
&cancellation,
)
});
thread::sleep(Duration::from_millis(100));
handle.cancel();
let error = request.join().unwrap().unwrap_err();
assert_eq!(error.to_string(), "prompt canceled");
});
assert!(pending.lock().unwrap().is_empty());
thread::sleep(Duration::from_millis(1100));
assert!(transport.lock().unwrap().is_process_running());
transport
.lock()
.unwrap()
.request(
"initialize",
json!({"rootUri":"file:///tmp"}),
Duration::from_secs(1),
&AgentCancellation::default(),
)
.unwrap();
}
fn wait_for_notification(
transport: &LspStdioTransport,
timeout: Duration,
) -> Option<JsonRpcNotification> {
let started = Instant::now();
while started.elapsed() < timeout {
if let Some(notification) = transport.try_recv_notification() {
return Some(notification);
}
thread::sleep(Duration::from_millis(10));
}
None
}
}