fix(handler): do call handler methods when initialize server (#118)

This commit is contained in:
4t145 2025-04-10 15:06:45 +08:00 committed by GitHub
parent d0962ec9ee
commit 31da5ff02d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 40 additions and 14 deletions

View file

@ -504,15 +504,16 @@ where
T: IntoTransport<R, E, A>,
E: std::error::Error + Send + Sync + 'static,
{
serve_inner(service, transport, peer_info, Default::default(), ct).await
let (peer, peer_rx) = Peer::new(Arc::new(AtomicU32RequestIdProvider::default()), peer_info);
serve_inner(service, transport, peer, peer_rx, ct).await
}
#[instrument(skip_all)]
async fn serve_inner<R, S, T, E, A>(
mut service: S,
transport: T,
peer_info: R::PeerInfo,
id_provider: Arc<AtomicU32Provider>,
peer: Peer<R>,
mut peer_rx: tokio::sync::mpsc::Receiver<PeerSinkMessage<R>>,
ct: CancellationToken,
) -> Result<RunningService<R, S>, E>
where
@ -525,14 +526,13 @@ where
const SINK_PROXY_BUFFER_SIZE: usize = 64;
let (sink_proxy_tx, mut sink_proxy_rx) =
tokio::sync::mpsc::channel::<TxJsonRpcMessage<R>>(SINK_PROXY_BUFFER_SIZE);
let peer_info = peer.peer_info();
if R::IS_CLIENT {
tracing::info!(?peer_info, "Service initialized as client");
} else {
tracing::info!(?peer_info, "Service initialized as server");
}
let (peer, mut peer_proxy) = <Peer<R>>::new(id_provider, peer_info);
service.set_peer(peer.clone());
let mut local_responder_pool = HashMap::new();
let mut local_ct_pool = HashMap::<RequestId, CancellationToken>::new();
@ -576,7 +576,7 @@ where
break QuitReason::Closed
}
}
m = peer_proxy.recv() => {
m = peer_rx.recv() => {
if let Some(m) = m {
Event::ProxyMessage(m)
} else {

View file

@ -174,7 +174,8 @@ where
}),
);
sink.send(notification).await?;
serve_inner(service, (sink, stream), initialize_result, id_provider, ct).await
let (peer, peer_rx) = Peer::new(id_provider, initialize_result);
serve_inner(service, (sink, stream), peer, peer_rx, ct).await
}
macro_rules! method {

View file

@ -5,7 +5,7 @@ use super::*;
use crate::model::{
CancelledNotification, CancelledNotificationParam, ClientInfo, ClientJsonRpcMessage,
ClientNotification, ClientRequest, ClientResult, CreateMessageRequest,
CreateMessageRequestParam, CreateMessageResult, ListRootsRequest, ListRootsResult,
CreateMessageRequestParam, CreateMessageResult, ErrorData, ListRootsRequest, ListRootsResult,
LoggingMessageNotification, LoggingMessageNotificationParam, ProgressNotification,
ProgressNotificationParam, PromptListChangedNotification, ResourceListChangedNotification,
ResourceUpdatedNotification, ResourceUpdatedNotificationParam, ServerInfo, ServerNotification,
@ -41,6 +41,12 @@ pub enum ServerError {
#[error("connection closed: {0}")]
ConnectionClosed(String),
#[error("unexpected initialize result: {0:?}")]
UnexpectedInitializeResponse(ServerResult),
#[error("initialize failed: {0}")]
InitializeFailed(ErrorData),
#[error("IO error: {0}")]
Io(#[from] std::io::Error),
}
@ -144,14 +150,34 @@ where
.await
.map_err(handle_server_error)?;
let ClientRequest::InitializeRequest(peer_info) = request else {
let ClientRequest::InitializeRequest(peer_info) = &request else {
return Err(handle_server_error(ServerError::ExpectedInitRequest(Some(
ClientJsonRpcMessage::request(request, id),
))));
};
let (peer, peer_rx) = Peer::new(id_provider, peer_info.params.clone());
let context = RequestContext {
ct: ct.child_token(),
id: id.clone(),
meta: request.get_meta().clone(),
extensions: request.extensions().clone(),
peer: peer.clone(),
};
// Send initialize response
let mut init_response = service.get_info();
let init_response = service.handle_request(request.clone(), context).await;
let mut init_response = match init_response {
Ok(ServerResult::InitializeResult(init_response)) => init_response,
Ok(result) => {
return Err(handle_server_error(
ServerError::UnexpectedInitializeResponse(result),
));
}
Err(e) => {
sink.send(ServerJsonRpcMessage::error(e.clone(), id))
.await?;
return Err(handle_server_error(ServerError::InitializeFailed(e)));
}
};
let protocol_version = match peer_info
.params
.protocol_version
@ -174,15 +200,14 @@ where
let notification = expect_notification(&mut stream, "initialize notification")
.await
.map_err(handle_server_error)?;
let ClientNotification::InitializedNotification(_) = notification else {
return Err(handle_server_error(ServerError::ExpectedInitNotification(
Some(ClientJsonRpcMessage::notification(notification)),
)));
};
let _ = service.handle_notification(notification).await;
// Continue processing service
serve_inner(service, (sink, stream), peer_info.params, id_provider, ct).await
serve_inner(service, (sink, stream), peer, peer_rx, ct).await
}
macro_rules! method {