use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use lsp_types::error_codes::SERVER_CANCELLED;
use lsp_types::notification::{Exit, Initialized, Notification};
use lsp_types::request::{Initialize, Request, Shutdown};
use lsp_types::{
ClientCapabilities, ClientInfo, InitializeParams, InitializedParams, WorkDoneProgressParams,
};
use serde_json::Value;
use tokio_util::sync::CancellationToken;
use tracing::{Instrument, debug};
use crate::builder::SharedHandler;
use crate::client::ClientHandle;
use crate::codec::{decode_value, encode_body, erase_value};
use crate::error::{BuildError, LspError};
use crate::failure::FailureReporter;
use crate::raw::{RawMessage, RequestId};
use crate::resource_policy::ResourcePolicy;
use crate::runtime::{Runtime, TaskFuture, TaskSend, default_runtime, ensure_runtime_available};
use crate::service::{HandlerTimeout, ServiceResult};
use crate::session::{
INBOUND_CAPACITY_EXHAUSTED, InboundReserveError, ProtocolControl, ProtocolSession,
SessionInput, run_handler_with_deadline,
};
use crate::telemetry::{ConnectionTrace, Direction};
use crate::transport::{Transport, TransportError, TransportReader};
use crate::{ClientError, Error, Outcome, Result};
type ReverseFuture = Pin<Box<dyn TaskFuture<std::result::Result<Value, LspError>>>>;
#[cfg(not(target_arch = "wasm32"))]
type ReverseHandler = Arc<dyn Fn(Value, CancellationToken) -> ReverseFuture + Send + Sync>;
#[cfg(target_arch = "wasm32")]
type ReverseHandler = Arc<dyn Fn(Value, CancellationToken) -> ReverseFuture>;
type ConnectionFuture = Pin<Box<dyn TaskFuture<Result<Outcome>>>>;
pub struct Client<T = ()> {
transport: T,
capabilities: ClientCapabilities,
client_info: Option<ClientInfo>,
initialization_options: Option<Value>,
handlers: HashMap<&'static str, ReverseHandler>,
resource_policy: ResourcePolicy,
}
impl Client<()> {
pub fn builder(capabilities: ClientCapabilities) -> ClientBuilder {
ClientBuilder::new(capabilities)
}
}
impl<T: Transport> Client<T> {
pub async fn connect(self) -> Result<ClientConnection> {
ensure_runtime_available()?;
let trace = ConnectionTrace::new();
let failure_reporter = FailureReporter::new(None, trace.id());
let span = trace.span();
let (mut reader, writer) = self.transport.split();
let (protocol, peer) = ProtocolSession::start(
default_runtime(),
self.resource_policy,
writer,
trace,
span.clone(),
failure_reporter.clone(),
ClientCloseCause::writer_failed,
ClientHandle::new,
);
let lifecycle = Arc::new(ClientLifecycle {
phase: Mutex::new(ClientPhase::Initializing),
protocol: protocol.control(),
});
let server = ServerHandle {
inner: peer.clone(),
lifecycle: Arc::clone(&lifecycle),
};
let mut engine = ClientEngine {
handlers: self.handlers,
protocol,
peer,
trace,
lifecycle,
};
let params = initialize_params(
self.capabilities,
self.client_info,
self.initialization_options,
);
if let Err(error) = engine.initialize(&mut reader, params).await {
engine
.protocol
.request_close(ClientCloseCause::InitializeFailed);
engine.protocol.close().await;
trace.connection_closed("initialize_failed");
return Err(error);
}
if let Err(error) = server
.inner
.notify_required::<Initialized>(InitializedParams {})
{
server
.lifecycle
.protocol
.request_close(ClientCloseCause::InitializeFailed);
engine.protocol.close().await;
trace.connection_closed("initialize_failed");
return Err(error.into());
}
server.lifecycle.mark_running();
Ok(ClientConnection {
server,
driver: Box::pin(engine.serve(reader).instrument(span)),
})
}
}
pub struct ClientBuilder {
capabilities: ClientCapabilities,
client_info: Option<ClientInfo>,
initialization_options: Option<Value>,
handlers: HashMap<&'static str, ReverseHandler>,
resource_policy: ResourcePolicy,
error: Option<BuildError>,
}
impl ClientBuilder {
fn new(capabilities: ClientCapabilities) -> Self {
Self {
capabilities,
client_info: None,
initialization_options: None,
handlers: HashMap::new(),
resource_policy: ResourcePolicy::default(),
error: None,
}
}
pub fn client_info(mut self, client_info: ClientInfo) -> Self {
self.client_info = Some(client_info);
self
}
pub fn initialization_options(mut self, options: Value) -> Self {
self.initialization_options = Some(options);
self
}
pub fn resource_policy(mut self, policy: ResourcePolicy) -> Self {
self.resource_policy = policy;
self
}
pub fn request<R, H, Fut>(mut self, handler: H) -> Self
where
R: Request,
H: Fn(R::Params, CancellationToken) -> Fut
+ SharedHandler<(R::Params, CancellationToken), Fut>
+ 'static,
Fut: Future<Output = std::result::Result<R::Result, LspError>> + TaskSend + 'static,
{
if is_reserved_reverse_method(R::METHOD) {
self.record(BuildError::ReservedMethod(R::METHOD.to_string()));
return self;
}
if self.handlers.contains_key(R::METHOD) {
self.record(BuildError::DuplicateMethod(R::METHOD.to_string()));
return self;
}
let handler = Arc::new(handler);
self.handlers.insert(
R::METHOD,
Arc::new(move |params, cancellation| {
let handler = Arc::clone(&handler);
let params =
serde_json::from_value::<R::Params>(params).map_err(LspError::invalid_params);
Box::pin(async move {
let params = params?;
let result = handler.invoke((params, cancellation)).await?;
erase_value(result)
})
}),
);
self
}
pub fn build<T: Transport>(
mut self,
transport: T,
) -> std::result::Result<Client<T>, BuildError> {
if let Err(error) = self.resource_policy.validate() {
self.record(error);
}
if let Some(error) = self.error {
return Err(error);
}
Ok(Client {
transport,
capabilities: self.capabilities,
client_info: self.client_info,
initialization_options: self.initialization_options,
handlers: self.handlers,
resource_policy: self.resource_policy,
})
}
fn record(&mut self, error: BuildError) {
if self.error.is_none() {
self.error = Some(error);
}
}
}
pub struct ClientConnection {
server: ServerHandle,
driver: ConnectionFuture,
}
impl ClientConnection {
pub fn server(&self) -> ServerHandle {
self.server.clone()
}
pub async fn serve(self) -> Result<Outcome> {
self.driver.await
}
}
#[derive(Clone)]
pub struct ServerHandle {
inner: ClientHandle,
lifecycle: Arc<ClientLifecycle>,
}
impl std::fmt::Debug for ServerHandle {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ServerHandle")
.finish_non_exhaustive()
}
}
impl ServerHandle {
pub fn notify<N>(&self, params: N::Params) -> std::result::Result<(), ClientError>
where
N: Notification,
{
self.lifecycle.ensure_running(N::METHOD)?;
self.inner.notify::<N>(params)
}
pub async fn request<R>(&self, params: R::Params) -> std::result::Result<R::Result, ClientError>
where
R: Request,
{
self.lifecycle.ensure_running(R::METHOD)?;
self.inner.request::<R>(params).await
}
pub async fn shutdown(&self) -> std::result::Result<(), ClientError> {
self.lifecycle.begin_shutdown()?;
let result = self.inner.request::<Shutdown>(()).await;
self.lifecycle.finish_shutdown(&result);
result
}
pub fn exit(&self) -> std::result::Result<(), ClientError> {
self.lifecycle.begin_exit()?;
let result = self.inner.notify_required::<Exit>(());
self.lifecycle
.protocol
.request_close(ClientCloseCause::Exit);
result
}
pub fn disconnect(&self) {
if self.lifecycle.disconnect() {
self.lifecycle
.protocol
.request_close(ClientCloseCause::Disconnect);
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ClientPhase {
Initializing,
Running,
ShutdownPending,
ShuttingDown,
Exited,
Disconnected,
}
#[derive(Clone, Copy)]
enum ClientTransition {
Initialized,
BeginShutdown,
ShutdownSucceeded,
ShutdownFailed,
Exit,
Disconnect,
}
#[derive(Clone, Copy)]
enum ClientWork {
Outbound,
Reverse,
}
impl ClientPhase {
fn as_str(self) -> &'static str {
match self {
Self::Initializing => "initializing",
Self::Running => "running",
Self::ShutdownPending => "shutdown is pending",
Self::ShuttingDown => "shut down",
Self::Exited => "exited",
Self::Disconnected => "disconnected",
}
}
fn transition(self, transition: ClientTransition) -> Option<Self> {
match (self, transition) {
(Self::Initializing, ClientTransition::Initialized) => Some(Self::Running),
(Self::Running, ClientTransition::BeginShutdown) => Some(Self::ShutdownPending),
(Self::ShutdownPending, ClientTransition::ShutdownSucceeded) => {
Some(Self::ShuttingDown)
}
(Self::ShutdownPending, ClientTransition::ShutdownFailed) => Some(Self::Running),
(Self::ShuttingDown, ClientTransition::Exit) => Some(Self::Exited),
(Self::Exited | Self::Disconnected, ClientTransition::Disconnect) => None,
(_, ClientTransition::Disconnect) => Some(Self::Disconnected),
_ => None,
}
}
fn permits(self, work: ClientWork) -> bool {
match work {
ClientWork::Outbound => self == Self::Running,
ClientWork::Reverse => matches!(self, Self::Running | Self::ShutdownPending),
}
}
}
struct ClientLifecycle {
phase: Mutex<ClientPhase>,
protocol: ProtocolControl<ClientHandle, ClientCloseCause>,
}
impl ClientLifecycle {
fn mark_running(&self) {
self.transition("initialize", ClientTransition::Initialized)
.expect("initialization completes only from the initializing phase");
}
fn ensure_running(&self, operation: &'static str) -> std::result::Result<(), ClientError> {
let phase = *self.phase.lock().unwrap();
if phase.permits(ClientWork::Outbound) && !is_reserved_reverse_method(operation) {
Ok(())
} else {
Err(Self::invalid(operation, phase))
}
}
fn begin_shutdown(&self) -> std::result::Result<(), ClientError> {
self.transition("shutdown", ClientTransition::BeginShutdown)
}
fn finish_shutdown(&self, result: &std::result::Result<(), ClientError>) {
let transition = if result.is_ok() {
ClientTransition::ShutdownSucceeded
} else {
ClientTransition::ShutdownFailed
};
if self.transition("shutdown", transition).is_ok() && result.is_ok() {
self.protocol.successful_shutdown();
}
}
fn begin_exit(&self) -> std::result::Result<(), ClientError> {
self.transition("exit", ClientTransition::Exit)
}
fn disconnect(&self) -> bool {
self.transition("disconnect", ClientTransition::Disconnect)
.is_ok()
}
fn rejects_reverse_work(&self) -> bool {
!self.phase.lock().unwrap().permits(ClientWork::Reverse)
}
fn transition(
&self,
operation: &'static str,
transition: ClientTransition,
) -> std::result::Result<(), ClientError> {
let mut phase = self.phase.lock().unwrap();
let current = *phase;
if let Some(next) = current.transition(transition) {
*phase = next;
Ok(())
} else {
Err(Self::invalid(operation, current))
}
}
fn invalid(operation: &'static str, phase: ClientPhase) -> ClientError {
ClientError::InvalidLifecycle {
operation,
state: phase.as_str(),
}
}
}
#[derive(Debug)]
enum ClientCloseCause {
Exit,
Disconnect,
ReaderEof,
ReaderFailed(TransportError),
WriterFailed,
InitializeFailed,
}
impl ClientCloseCause {
fn writer_failed() -> Self {
Self::WriterFailed
}
fn as_str(&self) -> &'static str {
match self {
Self::Exit => "exit",
Self::Disconnect => "disconnect",
Self::ReaderEof => "reader_eof",
Self::ReaderFailed(_) => "reader_failed",
Self::WriterFailed => "writer_failed",
Self::InitializeFailed => "initialize_failed",
}
}
fn into_result(self) -> Result<Outcome> {
match self {
Self::Exit => Ok(Outcome::Exit { code: 0 }),
Self::Disconnect => Ok(Outcome::TransportClosed),
Self::ReaderEof => Ok(Outcome::TransportClosed),
Self::ReaderFailed(error) => Err(Error::Transport(error)),
Self::WriterFailed => Ok(Outcome::WriterFailed),
Self::InitializeFailed => unreachable!("initialize failure returns from connect"),
}
}
}
struct ClientEngine<R> {
handlers: HashMap<&'static str, ReverseHandler>,
protocol: ProtocolSession<R, ClientHandle, ClientCloseCause>,
peer: ClientHandle,
trace: ConnectionTrace,
lifecycle: Arc<ClientLifecycle>,
}
impl<R: Runtime> ClientEngine<R> {
async fn initialize<Rd>(&mut self, reader: &mut Rd, params: InitializeParams) -> Result<()>
where
Rd: TransportReader,
{
let peer = self.peer.clone();
let mut pending = Box::pin(peer.request::<Initialize>(params));
loop {
match futures_util::future::select(pending, Box::pin(reader.recv())).await {
futures_util::future::Either::Left((result, _)) => {
result?;
return Ok(());
}
futures_util::future::Either::Right((message, still_pending)) => {
pending = still_pending;
match message {
Ok(message) => {
self.trace.message(Direction::Inbound, &message);
self.dispatch_during_initialize(message);
}
Err(error) => return Err(Error::Transport(error)),
}
}
}
}
}
fn dispatch_during_initialize(&self, message: RawMessage) {
match message {
RawMessage::Response { id, result } => self.complete_response(id, result),
RawMessage::Request { id, method, .. } if method == "initialize" => self
.protocol
.reject_inbound(id, LspError::invalid_request("duplicate initialize")),
RawMessage::Request { id, .. } => self
.protocol
.reject_inbound(id, LspError::ServerNotInitialized),
RawMessage::Notification { method, .. } => {
debug!(%method, "notification during client initialization ignored");
}
RawMessage::ProtocolError { error } => self.protocol.send_protocol_error(error),
}
}
async fn serve<Rd>(mut self, mut reader: Rd) -> Result<Outcome>
where
Rd: TransportReader,
{
loop {
let message = match self.protocol.next_input(&mut reader).await {
SessionInput::CloseRequested => break,
SessionInput::OutboundFailed => {
self.protocol.request_close(ClientCloseCause::WriterFailed);
break;
}
SessionInput::Message(message) => message,
};
match message {
Ok(message) => {
self.trace.message(Direction::Inbound, &message);
self.dispatch(message);
}
Err(TransportError::Closed) => {
self.protocol.request_close(ClientCloseCause::ReaderEof);
break;
}
Err(error) => {
self.protocol
.request_close(ClientCloseCause::ReaderFailed(error));
break;
}
}
}
self.lifecycle.disconnect();
self.protocol.close().await;
let cause = self.protocol.final_close_cause();
self.trace.connection_closed(cause.as_str());
cause.into_result()
}
fn dispatch(&mut self, message: RawMessage) {
match message {
RawMessage::Response { id, result } => self.complete_response(id, result),
RawMessage::Request { id, method, params } => {
self.dispatch_request(id, method.into_owned(), params)
}
RawMessage::Notification { method, params } if method == "$/cancelRequest" => {
#[derive(serde::Deserialize)]
struct CancelParams {
id: RequestId,
}
match crate::codec::decode_params::<CancelParams>(¶ms) {
Ok(params) => {
if let Some(reservation) = self.protocol.cancel_inbound(¶ms.id) {
self.protocol.complete_cancelled(reservation);
}
}
Err(error) => debug!(%error, "malformed cancellation ignored"),
}
}
RawMessage::Notification { method, .. } => {
debug!(%method, "unregistered server notification ignored");
}
RawMessage::ProtocolError { error } => self.protocol.send_protocol_error(error),
}
}
fn complete_response(
&self,
id: RequestId,
result: std::result::Result<bytes::Bytes, crate::JsonRpcError>,
) {
let delivered = match id {
RequestId::Number(number) if number > 0 => self
.peer
.outbound_registry()
.complete(number as u32, result),
_ => false,
};
if !delivered {
debug!("ignoring response with unknown or non-numeric id");
}
}
fn dispatch_request(&mut self, id: RequestId, method: String, params: bytes::Bytes) {
let reserved = match self.protocol.reserve_inbound(id.clone(), &method, true) {
Ok(reserved) => reserved,
Err(InboundReserveError::DuplicateId) => {
self.protocol
.reject_inbound(id, LspError::invalid_request("duplicate request id"));
return;
}
Err(InboundReserveError::CapacityExhausted) => {
self.protocol.reject_inbound(
id,
LspError::ServerError {
code: SERVER_CANCELLED as i32,
message: INBOUND_CAPACITY_EXHAUSTED.to_string(),
data: None,
},
);
return;
}
};
let reservation = reserved.reservation;
let cancellation = reserved
.cancellation
.expect("reverse requests are cancellable");
if self.lifecycle.rejects_reverse_work() {
self.protocol.complete_inbound(
reservation,
Err(LspError::invalid_request("invalid request after shutdown")),
);
return;
}
if method == "initialize" {
self.protocol.complete_inbound(
reservation,
Err(LspError::invalid_request("duplicate initialize")),
);
return;
}
let Some(handler) = self.handlers.get(method.as_str()).cloned() else {
self.protocol.complete_inbound(
reservation,
Err(LspError::MethodNotFound(method.to_string())),
);
return;
};
let params = match decode_value(¶ms) {
Ok(params) => params,
Err(error) => {
self.protocol.complete_inbound(reservation, Err(error));
return;
}
};
let completion = self.protocol.completion_gate();
let permit = Arc::clone(&reservation._permit);
let timeout = HandlerTimeout::new(
self.protocol.handler_timeout(),
self.trace,
method.clone(),
reservation.id.clone(),
);
timeout.arm();
let span = self.trace.request_span(&method, &reservation.id);
self.protocol.spawn(
async move {
let handler_cancellation = cancellation.clone();
let result = run_handler_with_deadline(
async move {
match handler(params, handler_cancellation).await {
Ok(value) => ServiceResult::Response(value),
Err(error) => ServiceResult::Error(error),
}
},
cancellation,
timeout,
)
.await;
let result = match result {
ServiceResult::Response(value) => encode_body(&value),
ServiceResult::Error(error) => Err(error),
ServiceResult::NoResponse => {
Err(LspError::internal("reverse request returned no response"))
}
};
completion.complete(reservation, result);
}
.instrument(span),
permit,
);
}
}
fn is_reserved_reverse_method(method: &str) -> bool {
matches!(
method,
"initialize" | "initialized" | "shutdown" | "exit" | "$/cancelRequest"
)
}
#[allow(deprecated)]
fn initialize_params(
capabilities: ClientCapabilities,
client_info: Option<ClientInfo>,
initialization_options: Option<Value>,
) -> InitializeParams {
InitializeParams {
process_id: None,
root_path: None,
root_uri: None,
initialization_options,
capabilities,
trace: None,
workspace_folders: None,
client_info,
locale: None,
work_done_progress_params: WorkDoneProgressParams::default(),
}
}