use std::collections::HashMap;
use std::time::{Duration, Instant};
use crate::StdioTransport;
use crate::Transport;
use crate::codec::json_rpc::LspMessage;
use serde_json::Value;
#[derive(Default)]
pub struct InitializeParams {
pub root_uri: Option<String>,
}
pub struct LspSession {
transport: Box<dyn Transport>,
next_id: i64,
open_documents: HashMap<String, i32>,
last_used_at: Instant,
}
impl LspSession {
pub fn spawn(cmd: &str) -> Result<Self, anyhow::Error> {
Self::spawn_with_args(cmd, &[])
}
pub fn spawn_with_args(cmd: &str, extra_args: &[String]) -> Result<Self, anyhow::Error> {
let transport = StdioTransport::spawn(cmd, extra_args)?;
Ok(Self {
transport: Box::new(transport),
next_id: 1,
open_documents: HashMap::new(),
last_used_at: Instant::now(),
})
}
pub fn with_transport(transport: Box<dyn Transport>) -> Self {
Self {
transport,
next_id: 1,
open_documents: HashMap::new(),
last_used_at: Instant::now(),
}
}
fn touch(&mut self) {
self.last_used_at = Instant::now();
}
pub fn last_used_at(&self) -> Instant {
self.last_used_at
}
pub async fn initialize(&mut self, params: InitializeParams) -> Result<Value, anyhow::Error> {
self.touch();
let root_uri = params.root_uri.map(|p| {
if p.starts_with("file://") {
p
} else {
format!("file://{p}")
}
});
let workspace_folders: Option<Vec<Value>> = root_uri.as_ref().map(|uri| {
let name = uri
.rsplit('/')
.next()
.filter(|s| !s.is_empty())
.unwrap_or("workspace");
vec![serde_json::json!({ "uri": uri, "name": name })]
});
let init_params = serde_json::json!({
"processId": null,
"capabilities": {},
"rootUri": root_uri,
"workspaceFolders": workspace_folders,
"clientInfo": {
"name": "lspz",
"version": env!("CARGO_PKG_VERSION"),
},
});
let result = self.send_request("initialize", init_params).await?;
self.send_notification("initialized", serde_json::json!({}))
.await?;
Ok(result)
}
pub async fn send_request(
&mut self,
method: &str,
params: Value,
) -> Result<Value, anyhow::Error> {
self.touch();
let id = self.next_id;
self.next_id += 1;
let msg = LspMessage::Request {
id,
method: method.into(),
params,
};
let frame = msg.to_bytes()?;
self.transport.send(&frame).await?;
loop {
let raw = tokio::time::timeout(Duration::from_secs(30), self.transport.receive())
.await
.map_err(|_| anyhow::anyhow!("timeout waiting for response to '{method}'"))??;
self.touch();
let parsed = LspMessage::from_frame_bytes(&raw)?;
match parsed {
LspMessage::Response {
id: rid,
result,
error,
} if rid == id => {
if let Some(err) = error {
anyhow::bail!("LSP error {}: {}", err.code, err.message);
}
return Ok(result.unwrap_or(Value::Null));
}
LspMessage::Response {
id: rid, ref error, ..
} => {
tracing::warn!(
"Ignoring response for id {} (waiting for {}): {:?}",
rid,
id,
error
);
}
LspMessage::Notification { method: m, .. } => {
tracing::trace!("Buffered notification: {}", m);
}
LspMessage::Request { method: m, .. } => {
tracing::trace!("Ignored request during wait: {}", m);
}
}
}
}
pub async fn send_notification(
&mut self,
method: &str,
params: Value,
) -> Result<(), anyhow::Error> {
self.touch();
let msg = LspMessage::Notification {
method: method.into(),
params,
};
let frame = msg.to_bytes()?;
self.transport.send(&frame).await?;
Ok(())
}
pub async fn open_or_update_document(
&mut self,
uri: &str,
language_id: &str,
content: &str,
) -> Result<(), anyhow::Error> {
self.touch();
let new_version = if let Some(version) = self.open_documents.get_mut(uri) {
*version += 1;
Some(*version)
} else {
None
};
if let Some(version) = new_version {
self.send_notification(
"textDocument/didChange",
serde_json::json!({
"textDocument": { "uri": uri, "version": version },
"contentChanges": [{ "text": content }],
}),
)
.await?;
} else {
self.send_notification(
"textDocument/didOpen",
serde_json::json!({
"textDocument": {
"uri": uri,
"languageId": language_id,
"version": 1,
"text": content,
}
}),
)
.await?;
self.open_documents.insert(uri.to_string(), 1);
}
Ok(())
}
pub async fn wait_for_notification(&mut self, method: &str) -> Result<Value, anyhow::Error> {
self.wait_for_notification_where(method, |_| true).await
}
pub async fn wait_for_notification_where(
&mut self,
method: &str,
predicate: impl Fn(&Value) -> bool,
) -> Result<Value, anyhow::Error> {
self.touch();
loop {
let raw = tokio::time::timeout(Duration::from_secs(30), self.transport.receive())
.await
.map_err(|_| anyhow::anyhow!("timeout waiting for '{method}' notification"))??;
self.touch();
let parsed = LspMessage::from_frame_bytes(&raw)?;
match parsed {
LspMessage::Notification { method: m, params }
if m == method && predicate(¶ms) =>
{
return Ok(params);
}
LspMessage::Notification { method: m, .. } => {
tracing::trace!("Skipping notification: {}", m);
}
LspMessage::Response { id, ref result, .. } => {
tracing::trace!("Skipping response id={}: {:?}", id, result);
}
LspMessage::Request { method: m, .. } => {
tracing::trace!("Skipping request: {}", m);
}
}
}
}
pub fn try_wait(&mut self) -> Result<Option<std::process::ExitStatus>, anyhow::Error> {
Ok(self.transport.try_wait()?)
}
#[cfg(test)]
pub fn mock_sent_messages(&mut self) -> Vec<Vec<u8>> {
use crate::transport::mock::MockTransport;
self.transport
.as_any_mut()
.and_then(|any| any.downcast_mut::<MockTransport>())
.map(|m| m.sent_messages())
.unwrap_or_default()
}
}
#[cfg(test)]
mod tests {
use crate::codec::json_rpc::LspMessage;
use crate::transport::mock::MockTransport;
use serde_json::json;
use super::*;
#[tokio::test]
async fn test_initialize_default_null_root() {
let mock = MockTransport::new();
mock.push_message(&LspMessage::Response {
id: 1,
result: Some(json!({ "capabilities": {} })),
error: None,
})
.unwrap();
let mut session = LspSession::with_transport(Box::new(mock));
session
.initialize(InitializeParams::default())
.await
.unwrap();
let sent = session.mock_sent_messages();
assert_eq!(sent.len(), 2);
let init_msg = LspMessage::from_frame_bytes(&sent[0]).unwrap();
if let LspMessage::Request { params, .. } = init_msg {
assert_eq!(params["rootUri"], serde_json::Value::Null);
} else {
panic!("Expected request, got {init_msg:?}");
}
}
#[tokio::test]
async fn test_initialize_with_root_uri() {
let mock = MockTransport::new();
mock.push_message(&LspMessage::Response {
id: 1,
result: Some(json!({ "capabilities": {} })),
error: None,
})
.unwrap();
let mut session = LspSession::with_transport(Box::new(mock));
session
.initialize(InitializeParams {
root_uri: Some("/home/user/project".into()),
})
.await
.unwrap();
let sent = session.mock_sent_messages();
let init_msg = LspMessage::from_frame_bytes(&sent[0]).unwrap();
if let LspMessage::Request { params, .. } = init_msg {
assert_eq!(params["rootUri"], "file:///home/user/project");
} else {
panic!("Expected request, got {init_msg:?}");
}
}
#[tokio::test]
async fn test_initialize_root_uri_already_prefixed() {
let mock = MockTransport::new();
mock.push_message(&LspMessage::Response {
id: 1,
result: Some(json!({ "capabilities": {} })),
error: None,
})
.unwrap();
let mut session = LspSession::with_transport(Box::new(mock));
session
.initialize(InitializeParams {
root_uri: Some("file:///home/user/project".into()),
})
.await
.unwrap();
let sent = session.mock_sent_messages();
let init_msg = LspMessage::from_frame_bytes(&sent[0]).unwrap();
if let LspMessage::Request { params, .. } = init_msg {
assert_eq!(params["rootUri"], "file:///home/user/project");
} else {
panic!("Expected request, got {init_msg:?}");
}
}
#[tokio::test]
async fn test_open_or_update_first_call_sends_did_open() {
let mock = MockTransport::new();
let mut session = LspSession::with_transport(Box::new(mock));
session
.open_or_update_document("file:///test.py", "python", "print('hi')")
.await
.unwrap();
let sent = session.mock_sent_messages();
assert_eq!(sent.len(), 1);
let msg = LspMessage::from_frame_bytes(&sent[0]).unwrap();
if let LspMessage::Notification { method, params } = msg {
assert_eq!(method, "textDocument/didOpen");
assert_eq!(params["textDocument"]["uri"], "file:///test.py");
assert_eq!(params["textDocument"]["version"], 1);
assert_eq!(params["textDocument"]["languageId"], "python");
} else {
panic!("Expected notification, got {msg:?}");
}
}
#[tokio::test]
async fn test_open_or_update_second_call_sends_did_change() {
let mock = MockTransport::new();
let mut session = LspSession::with_transport(Box::new(mock));
session
.open_or_update_document("file:///test.py", "python", "v1")
.await
.unwrap();
session
.open_or_update_document("file:///test.py", "python", "v2")
.await
.unwrap();
let sent = session.mock_sent_messages();
assert_eq!(sent.len(), 2);
let msg1 = LspMessage::from_frame_bytes(&sent[0]).unwrap();
if let LspMessage::Notification { method, .. } = msg1 {
assert_eq!(method, "textDocument/didOpen");
} else {
panic!("Expected notification, got {msg1:?}");
}
let msg2 = LspMessage::from_frame_bytes(&sent[1]).unwrap();
if let LspMessage::Notification { method, params } = msg2 {
assert_eq!(method, "textDocument/didChange");
assert_eq!(params["textDocument"]["version"], 2);
assert_eq!(params["contentChanges"][0]["text"], "v2");
} else {
panic!("Expected notification, got {msg2:?}");
}
}
#[tokio::test]
async fn test_open_or_update_increments_version() {
let mock = MockTransport::new();
let mut session = LspSession::with_transport(Box::new(mock));
session
.open_or_update_document("file:///a.rs", "rust", "1")
.await
.unwrap();
session
.open_or_update_document("file:///a.rs", "rust", "2")
.await
.unwrap();
session
.open_or_update_document("file:///a.rs", "rust", "3")
.await
.unwrap();
let sent = session.mock_sent_messages();
assert_eq!(sent.len(), 3);
let msg0 = LspMessage::from_frame_bytes(&sent[0]).unwrap();
if let LspMessage::Notification { method, params } = msg0 {
assert_eq!(method, "textDocument/didOpen");
assert_eq!(params["textDocument"]["version"], 1);
}
let msg1 = LspMessage::from_frame_bytes(&sent[1]).unwrap();
if let LspMessage::Notification { method, params } = msg1 {
assert_eq!(method, "textDocument/didChange");
assert_eq!(params["textDocument"]["version"], 2);
}
let msg2 = LspMessage::from_frame_bytes(&sent[2]).unwrap();
if let LspMessage::Notification { method, params } = msg2 {
assert_eq!(method, "textDocument/didChange");
assert_eq!(params["textDocument"]["version"], 3);
}
}
#[tokio::test]
async fn test_open_or_update_independent_documents() {
let mock = MockTransport::new();
let mut session = LspSession::with_transport(Box::new(mock));
session
.open_or_update_document("file:///a.py", "python", "a")
.await
.unwrap();
session
.open_or_update_document("file:///b.py", "python", "b")
.await
.unwrap();
session
.open_or_update_document("file:///a.py", "python", "a2")
.await
.unwrap();
let sent = session.mock_sent_messages();
assert_eq!(sent.len(), 3);
let msg0 = LspMessage::from_frame_bytes(&sent[0]).unwrap();
if let LspMessage::Notification { method, params } = msg0 {
assert_eq!(method, "textDocument/didOpen");
assert_eq!(params["textDocument"]["uri"], "file:///a.py");
}
let msg1 = LspMessage::from_frame_bytes(&sent[1]).unwrap();
if let LspMessage::Notification { method, params } = msg1 {
assert_eq!(method, "textDocument/didOpen");
assert_eq!(params["textDocument"]["uri"], "file:///b.py");
}
let msg2 = LspMessage::from_frame_bytes(&sent[2]).unwrap();
if let LspMessage::Notification { method, params } = msg2 {
assert_eq!(method, "textDocument/didChange");
assert_eq!(params["textDocument"]["uri"], "file:///a.py");
assert_eq!(params["textDocument"]["version"], 2);
}
}
}