fix: inject part into extension when handing init req (#275)

This commit is contained in:
4t145 2025-06-24 00:52:38 +08:00 committed by GitHub
parent a62c6d1db8
commit 6d3190504c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 22 additions and 3 deletions

View file

@ -103,6 +103,15 @@ pub(crate) const fn internal_error_response<E: Display>(
}
}
pub(crate) fn unexpected_message_response(
expect: &str,
) -> Response<UnsyncBoxBody<Bytes, Infallible>> {
Response::builder()
.status(http::StatusCode::UNPROCESSABLE_ENTITY)
.body(Full::new(Bytes::from(format!("Unexpected message, expect {expect}"))).boxed_unsync())
.expect("valid response")
}
pub(crate) async fn expect_json<B>(
body: B,
) -> Result<ClientJsonRpcMessage, Response<UnsyncBoxBody<Bytes, Infallible>>>

View file

@ -10,7 +10,7 @@ use tokio_stream::wrappers::ReceiverStream;
use super::session::SessionManager;
use crate::{
RoleServer,
model::{ClientJsonRpcMessage, GetExtensions},
model::{ClientJsonRpcMessage, ClientRequest, GetExtensions},
serve_server,
service::serve_directly,
transport::{
@ -21,7 +21,7 @@ use crate::{
},
server_side_http::{
BoxResponse, ServerSseMessage, accepted_response, expect_json,
internal_error_response, sse_stream_response,
internal_error_response, sse_stream_response, unexpected_message_response,
},
},
},
@ -318,6 +318,15 @@ where
.create_session()
.await
.map_err(internal_error_response("create session"))?;
if let ClientJsonRpcMessage::Request(req) = &mut message {
if !matches!(req.request, ClientRequest::InitializeRequest(_)) {
return Err(unexpected_message_response("initialize request"));
}
// inject request part to extensions
req.request.extensions_mut().insert(part);
} else {
return Err(unexpected_message_response("initialize request"));
}
let service = self
.get_service()
.map_err(internal_error_response("get service"))?;
@ -378,7 +387,8 @@ where
.get_service()
.map_err(internal_error_response("get service"))?;
match message {
ClientJsonRpcMessage::Request(request) => {
ClientJsonRpcMessage::Request(mut request) => {
request.request.extensions_mut().insert(part);
let (transport, receiver) =
OneshotTransport::<RoleServer>::new(ClientJsonRpcMessage::Request(request));
let service = serve_directly(service, transport, None);