use std::collections::HashMap;
use std::sync::Arc;
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::select;
use tokio::sync::RwLock;
use crate::codec::compact::CompactDiagnostics;
use crate::codec::json_rpc;
use crate::codec::toon;
use crate::config::{Config, OutputFormat};
use crate::error::LspzError;
use crate::interceptors::workspace_diagnostics::workspace_diagnostics_to_toon;
use crate::interceptors::workspace_symbols::workspace_symbols_to_toon;
use crate::interceptors::{Direction, InterceptorChain};
use crate::transport::Transport;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum State {
Created,
Initializing,
Ready,
ShuttingDown,
Exited,
}
pub struct Proxy {
config: Arc<RwLock<Config>>,
state: State,
transport: Box<dyn Transport>,
interceptor_chain: InterceptorChain,
pending_requests: HashMap<u64, String>,
_config_watcher: Option<crate::config_watcher::ConfigWatcher>,
}
impl Proxy {
pub fn new(
config: Arc<RwLock<Config>>,
transport: Box<dyn Transport>,
interceptor_chain: InterceptorChain,
) -> Self {
Self {
config,
state: State::Created,
transport,
interceptor_chain,
pending_requests: HashMap::new(),
_config_watcher: None,
}
}
pub fn state(&self) -> State {
self.state
}
pub fn set_config_watcher(&mut self, watcher: crate::config_watcher::ConfigWatcher) {
self._config_watcher = Some(watcher);
}
pub fn shared_config(&self) -> Arc<RwLock<Config>> {
self.config.clone()
}
pub async fn start(&mut self) -> Result<(), LspzError> {
self.state = State::Initializing;
tracing::info!("Proxy starting (handshake)");
self.perform_handshake().await?;
self.state = State::Ready;
tracing::info!("Proxy ready, entering message loop");
self.message_loop().await
}
async fn perform_handshake(&mut self) -> Result<(), LspzError> {
let mut stdin = BufReader::new(tokio::io::stdin());
let mut stdout = tokio::io::stdout();
let init_req = read_stdin_frame(&mut stdin).await?;
ensure_method(&init_req, "initialize")?;
self.transport.send(&init_req).await?;
tracing::debug!("Forwarded 'initialize' request to server");
let init_resp = self.transport.receive().await?;
stdout.write_all(&init_resp).await?;
stdout.flush().await?;
tracing::debug!("Forwarded 'initialize' response to client");
let init_not = read_stdin_frame(&mut stdin).await?;
ensure_method(&init_not, "initialized")?;
self.transport.send(&init_not).await?;
tracing::debug!("Forwarded 'initialized' notification to server");
Ok(())
}
async fn message_loop(&mut self) -> Result<(), LspzError> {
let mut stdin = BufReader::new(tokio::io::stdin());
let mut stdout = tokio::io::stdout();
loop {
select! {
client_msg = read_stdin_frame(&mut stdin) => {
let msg_bytes = match client_msg {
Ok(bytes) => bytes,
Err(LspzError::ServerExited) => {
tracing::info!("Client stdin closed, shutting down");
self.state = State::Exited;
return Ok(());
}
Err(e) => {
tracing::error!(error = %e, "Error reading from client");
return Err(e);
}
};
if extract_method(&msg_bytes).is_ok_and(|m| m == "shutdown") {
tracing::info!("Received 'shutdown' from client");
self.transport.send(&msg_bytes).await?;
self.state = State::ShuttingDown;
let resp = self.transport.receive().await?;
stdout.write_all(&resp).await?;
stdout.flush().await?;
self.state = State::Exited;
return Ok(());
}
track_pending_request(&msg_bytes, &mut self.pending_requests);
self.transport.send(&msg_bytes).await?;
}
server_msg = self.transport.receive() => {
let msg_bytes = match server_msg {
Ok(bytes) => bytes,
Err(LspzError::ServerExited) => {
tracing::warn!("Server exited unexpectedly");
self.state = State::Exited;
return Err(LspzError::ServerExited);
}
Err(e) => {
tracing::error!(error = %e, "Error reading from server");
return Err(e);
}
};
let processed = self.process_server_message(&msg_bytes).await;
if !processed.is_empty() {
stdout.write_all(&processed).await?;
stdout.flush().await?;
}
}
}
}
}
async fn process_server_message(&mut self, raw: &[u8]) -> Vec<u8> {
let (frame, _) = match json_rpc::parse_frame(raw) {
Ok(Some(result)) => result,
_ => return raw.to_vec(),
};
let json_val: serde_json::Value = match serde_json::from_slice(&frame.body) {
Ok(v) => v,
Err(_) => return raw.to_vec(),
};
if let Some(method) = json_val.get("method").and_then(|v| v.as_str()) {
return self.process_notification(method, &json_val, raw).await;
}
if json_val.get("id").is_some() {
return self.process_response(&json_val, raw).await;
}
raw.to_vec()
}
async fn process_notification(
&self,
method: &str,
json_val: &serde_json::Value,
raw: &[u8],
) -> Vec<u8> {
let params = json_val
.get("params")
.cloned()
.unwrap_or(serde_json::Value::Null);
let transformed = match self
.interceptor_chain
.process(method, params, Direction::ServerToClient)
.await
{
Ok(Some(p)) => p,
Ok(None) => return Vec::new(), Err(_) => return raw.to_vec(), };
if self.config.read().await.output_format == OutputFormat::Toon {
return self.toon_output(method, &transformed, raw);
}
let mut obj = match json_val.as_object() {
Some(o) => o.clone(),
None => return raw.to_vec(),
};
obj.insert("params".into(), transformed);
let new_val = serde_json::Value::Object(obj);
match json_rpc::serialize_frame(&new_val) {
Ok(bytes) => bytes,
Err(_) => raw.to_vec(),
}
}
async fn process_response(&mut self, json_val: &serde_json::Value, raw: &[u8]) -> Vec<u8> {
let id = match json_val.get("id").and_then(|v| v.as_u64()) {
Some(id) => id,
None => return raw.to_vec(),
};
let method = match self.pending_requests.remove(&id) {
Some(m) => m,
None => return raw.to_vec(),
};
if json_val.get("error").is_some() {
return raw.to_vec();
}
let params = json_val
.get("result")
.cloned()
.unwrap_or(serde_json::Value::Null);
let transformed = match self
.interceptor_chain
.process(&method, params, Direction::ServerToClient)
.await
{
Ok(Some(p)) => p,
Ok(None) => return Vec::new(), Err(_) => return raw.to_vec(), };
let mut obj = match json_val.as_object() {
Some(o) => o.clone(),
None => return raw.to_vec(),
};
obj.insert("result".into(), transformed);
let new_val = serde_json::Value::Object(obj);
match json_rpc::serialize_frame(&new_val) {
Ok(bytes) => bytes,
Err(_) => raw.to_vec(),
}
}
fn toon_output(&self, method: &str, params: &serde_json::Value, raw: &[u8]) -> Vec<u8> {
let toon_text = match method {
"textDocument/publishDiagnostics" => {
match serde_json::from_value::<CompactDiagnostics>(params.clone()) {
Ok(compact) => toon::diagnostics_to_toon(&compact),
Err(_) => return raw.to_vec(),
}
}
"textDocument/completion" => match toon::completions_to_toon(params) {
Ok(t) => t,
Err(_) => return raw.to_vec(),
},
"textDocument/hover" => match toon::hover_to_toon(params) {
Ok(t) => t,
Err(_) => return raw.to_vec(),
},
"textDocument/documentSymbol" => match toon::symbols_to_toon(params) {
Ok(t) => t,
Err(_) => return raw.to_vec(),
},
"textDocument/references"
| "textDocument/definition"
| "textDocument/implementation"
| "textDocument/typeDefinition" => match toon::locations_to_toon(params) {
Ok(t) => t,
Err(_) => return raw.to_vec(),
},
"workspace/symbol" => match workspace_symbols_to_toon(params) {
Ok(t) => t,
Err(_) => return raw.to_vec(),
},
"workspace/diagnostic" => match workspace_diagnostics_to_toon(params) {
Ok(t) => t,
Err(_) => return raw.to_vec(),
},
_ => return raw.to_vec(),
};
let msg = serde_json::json!({
"jsonrpc": "2.0",
"method": method,
"params": {
"format": "toon",
"text": toon_text,
}
});
match json_rpc::serialize_frame(&msg) {
Ok(bytes) => bytes,
Err(_) => raw.to_vec(),
}
}
}
fn track_pending_request(raw: &[u8], pending: &mut HashMap<u64, String>) {
let (frame, _) = match json_rpc::parse_frame(raw) {
Ok(Some(f)) => f,
_ => return,
};
let val: serde_json::Value = match serde_json::from_slice(&frame.body) {
Ok(v) => v,
_ => return,
};
let id = match val.get("id").and_then(|v| v.as_u64()) {
Some(id) => id,
None => return,
};
let method = match val.get("method").and_then(|v| v.as_str()) {
Some(m) => m,
None => return,
};
pending.insert(id, method.to_string());
}
async fn read_stdin_frame(reader: &mut BufReader<tokio::io::Stdin>) -> Result<Vec<u8>, LspzError> {
let mut header = String::new();
loop {
let mut line = String::new();
let n = reader.read_line(&mut line).await.map_err(|e| {
if e.kind() == std::io::ErrorKind::UnexpectedEof {
LspzError::ServerExited
} else {
LspzError::Io(e)
}
})?;
if n == 0 {
return Err(LspzError::ServerExited);
}
header.push_str(&line);
if line == "\r\n" || line == "\n" {
break;
}
}
let content_length = crate::transport::framing::parse_content_length(&header)?;
let mut body = vec![0u8; content_length as usize];
reader.read_exact(&mut body).await.map_err(|e| {
if e.kind() == std::io::ErrorKind::UnexpectedEof {
LspzError::ServerExited
} else {
LspzError::Io(e)
}
})?;
Ok([header.as_bytes(), &body].concat())
}
fn extract_method(raw: &[u8]) -> Result<String, LspzError> {
let (frame, _) = json_rpc::parse_frame(raw)?
.ok_or_else(|| LspzError::Protocol("incomplete frame".into()))?;
let val: serde_json::Value = serde_json::from_slice(&frame.body)?;
val.get("method")
.and_then(|v| v.as_str())
.map(String::from)
.ok_or_else(|| LspzError::Protocol("no 'method' field in message".into()))
}
fn ensure_method(raw: &[u8], expected: &str) -> Result<(), LspzError> {
let method = extract_method(raw)?;
if method != expected {
return Err(LspzError::Protocol(format!(
"expected '{expected}', got '{method}'"
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::codec::json_rpc::LspMessage;
use crate::transport::framing::parse_content_length;
#[test]
fn test_extract_method() {
let msg = LspMessage::Notification {
method: "textDocument/publishDiagnostics".into(),
params: serde_json::Value::Null,
};
let bytes = msg.to_bytes().unwrap();
assert_eq!(
extract_method(&bytes).unwrap(),
"textDocument/publishDiagnostics"
);
}
#[test]
fn test_ensure_method_ok() {
let msg = LspMessage::Request {
id: 1,
method: "initialize".into(),
params: serde_json::Value::Null,
};
let bytes = msg.to_bytes().unwrap();
assert!(ensure_method(&bytes, "initialize").is_ok());
}
#[test]
fn test_ensure_method_err() {
let msg = LspMessage::Request {
id: 1,
method: "shutdown".into(),
params: serde_json::Value::Null,
};
let bytes = msg.to_bytes().unwrap();
assert!(ensure_method(&bytes, "initialize").is_err());
}
#[test]
fn test_parse_content_length() {
let header = "Content-Length: 42\r\n\r\n";
assert_eq!(parse_content_length(header).unwrap(), 42);
}
#[test]
fn test_missing_content_length() {
let header = "\r\n\r\n";
assert!(parse_content_length(header).is_err());
}
#[test]
fn test_invalid_content_length_value() {
let header = "Content-Length: abc\r\n\r\n";
assert!(parse_content_length(header).is_err());
}
}