use std::convert::Infallible;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use tracing::{debug, error, info, warn};
use bytes::Bytes;
use futures::Stream;
use http_body::{Body, Frame};
use http_body_util::{BodyExt, Full};
use hyper::header::{ACCEPT, CONTENT_TYPE};
use hyper::{Method, Request, Response, StatusCode};
use crate::middleware::bearer::{extract_bearer_token, is_bearer_scheme};
use chrono;
use turul_mcp_json_rpc_server::{
JsonRpcDispatcher,
r#async::SessionContext,
dispatch::{JsonRpcMessage, JsonRpcMessageResult, parse_json_rpc_message},
error::{JsonRpcError, JsonRpcErrorObject},
};
use turul_mcp_protocol::McpError;
use turul_mcp_protocol::ServerCapabilities;
use turul_mcp_session_storage::{InMemorySessionStorage, SessionView};
use uuid::Uuid;
use crate::{
Result, ServerConfig, StreamConfig, StreamManager,
json_rpc_responses::*,
notification_bridge::{SharedNotificationBroadcaster, StreamManagerNotificationBroadcaster},
protocol::{
extract_last_event_id, extract_protocol_version, extract_session_id, normalize_header_value,
},
};
use std::collections::HashMap;
pub struct SessionSseStream {
stream: Pin<Box<dyn Stream<Item = std::result::Result<Bytes, Infallible>> + Send>>,
}
impl SessionSseStream {
pub fn new<S>(stream: S) -> Self
where
S: Stream<Item = std::result::Result<Bytes, Infallible>> + Send + 'static,
{
Self {
stream: Box::pin(stream),
}
}
}
impl Drop for SessionSseStream {
fn drop(&mut self) {
debug!("DROP: SessionSseStream - HTTP response body being cleaned up");
debug!("This may indicate early cleanup of SSE response stream");
}
}
impl Body for SessionSseStream {
type Data = Bytes;
type Error = Infallible;
fn poll_frame(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<std::result::Result<Frame<Self::Data>, Self::Error>>> {
match self.stream.as_mut().poll_next(cx) {
Poll::Ready(Some(Ok(data))) => Poll::Ready(Some(Ok(Frame::data(data)))),
Poll::Ready(Some(Err(never))) => match never {},
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
}
}
}
type JsonRpcBody = Full<Bytes>;
type UnifiedMcpBody = http_body_util::combinators::UnsyncBoxBody<Bytes, hyper::Error>;
#[derive(Debug, Clone, PartialEq)]
enum AcceptMode {
Compliant,
JsonOnly,
SseOnly,
Invalid,
}
fn parse_mcp_accept_header(accept_header: &str) -> (AcceptMode, bool) {
let accepts_json = accept_header.contains("application/json") || accept_header.contains("*/*");
let accepts_sse = accept_header.contains("text/event-stream");
let mode = match (accepts_json, accepts_sse) {
(true, true) => AcceptMode::Compliant,
(true, false) => AcceptMode::JsonOnly, (false, true) => AcceptMode::SseOnly,
(false, false) => AcceptMode::Invalid,
};
let should_use_sse = match mode {
AcceptMode::Compliant => true, AcceptMode::JsonOnly => false, AcceptMode::SseOnly => true, AcceptMode::Invalid => false, };
(mode, should_use_sse)
}
fn convert_to_unified_body(full_body: Full<Bytes>) -> UnifiedMcpBody {
full_body.map_err(|never| match never {}).boxed_unsync()
}
fn jsonrpc_error_to_unified_body(error: JsonRpcError) -> Result<Response<UnifiedMcpBody>> {
let error_json = serde_json::to_string(&error)?;
Ok(Response::builder()
.status(StatusCode::OK) .header(CONTENT_TYPE, "application/json")
.body(convert_to_unified_body(Full::new(Bytes::from(error_json))))
.unwrap())
}
pub struct SessionMcpHandler {
pub(crate) config: ServerConfig,
pub(crate) dispatcher: Arc<JsonRpcDispatcher<McpError>>,
pub(crate) session_storage: Arc<turul_mcp_session_storage::BoxedSessionStorage>,
pub(crate) stream_config: StreamConfig,
pub(crate) stream_manager: Arc<StreamManager>,
pub(crate) middleware_stack: Arc<crate::middleware::MiddlewareStack>,
pub(crate) tool_fingerprint: Option<String>,
pub(crate) tool_notifier: Option<Arc<dyn crate::ToolChangeNotifier>>,
}
impl Clone for SessionMcpHandler {
fn clone(&self) -> Self {
Self {
config: self.config.clone(),
dispatcher: Arc::clone(&self.dispatcher),
session_storage: Arc::clone(&self.session_storage),
stream_config: self.stream_config.clone(),
stream_manager: Arc::clone(&self.stream_manager),
middleware_stack: Arc::clone(&self.middleware_stack),
tool_fingerprint: self.tool_fingerprint.clone(),
tool_notifier: self.tool_notifier.clone(),
}
}
}
impl SessionMcpHandler {
pub fn new(
config: ServerConfig,
dispatcher: Arc<JsonRpcDispatcher<McpError>>,
stream_config: StreamConfig,
) -> Self {
let storage: Arc<turul_mcp_session_storage::BoxedSessionStorage> =
Arc::new(InMemorySessionStorage::new());
let middleware_stack = Arc::new(crate::middleware::MiddlewareStack::new());
Self::with_storage(config, dispatcher, storage, stream_config, middleware_stack)
}
pub fn with_shared_stream_manager(
config: ServerConfig,
dispatcher: Arc<JsonRpcDispatcher<McpError>>,
session_storage: Arc<turul_mcp_session_storage::BoxedSessionStorage>,
stream_config: StreamConfig,
stream_manager: Arc<StreamManager>,
middleware_stack: Arc<crate::middleware::MiddlewareStack>,
) -> Self {
Self {
config,
dispatcher,
session_storage,
stream_config,
stream_manager,
middleware_stack,
tool_fingerprint: None,
tool_notifier: None,
}
}
pub fn with_storage(
config: ServerConfig,
dispatcher: Arc<JsonRpcDispatcher<McpError>>,
session_storage: Arc<turul_mcp_session_storage::BoxedSessionStorage>,
stream_config: StreamConfig,
middleware_stack: Arc<crate::middleware::MiddlewareStack>,
) -> Self {
let stream_manager = Arc::new(StreamManager::with_config(
Arc::clone(&session_storage),
stream_config.clone(),
));
Self {
config,
dispatcher,
session_storage,
stream_config,
stream_manager,
middleware_stack,
tool_fingerprint: None,
tool_notifier: None,
}
}
pub fn with_tool_fingerprint(mut self, fingerprint: Option<String>) -> Self {
self.tool_fingerprint = fingerprint;
self
}
pub fn with_tool_notifier(mut self, notifier: Arc<dyn crate::ToolChangeNotifier>) -> Self {
self.tool_notifier = Some(notifier);
self
}
pub fn get_stream_manager(&self) -> &Arc<StreamManager> {
&self.stream_manager
}
pub async fn handle_mcp_request<B>(&self, req: Request<B>) -> Result<Response<UnifiedMcpBody>>
where
B: http_body::Body<Data = bytes::Bytes, Error = hyper::Error> + Send + 'static,
{
debug!(
"SESSION HANDLER processing {} {}",
req.method(),
req.uri().path()
);
match *req.method() {
Method::POST => {
let response = self.handle_json_rpc_request(req).await?;
Ok(response)
}
Method::GET => self.handle_sse_request(req).await,
Method::DELETE => {
let response = self.handle_delete_request(req).await?;
Ok(response.map(convert_to_unified_body))
}
Method::OPTIONS => {
let response = self.handle_preflight();
Ok(response.map(convert_to_unified_body))
}
_ => {
let response = self.method_not_allowed();
Ok(response.map(convert_to_unified_body))
}
}
}
async fn handle_json_rpc_request<B>(&self, req: Request<B>) -> Result<Response<UnifiedMcpBody>>
where
B: http_body::Body<Data = bytes::Bytes, Error = hyper::Error> + Send + 'static,
{
let headers: HashMap<String, String> = req
.headers()
.iter()
.filter_map(|(k, v)| {
v.to_str()
.ok()
.map(|s| (k.as_str().to_string(), s.to_string()))
})
.collect();
let protocol_version = extract_protocol_version(req.headers());
let session_id = extract_session_id(req.headers());
debug!(
"POST request - Protocol: {}, Session: {:?}",
protocol_version, session_id
);
let content_type = req
.headers()
.get(CONTENT_TYPE)
.and_then(|ct| ct.to_str().ok())
.map(normalize_header_value)
.unwrap_or_default();
if !content_type.starts_with("application/json") {
warn!("Invalid content type: {}", content_type);
return Ok(
bad_request_response("Content-Type must be application/json")
.map(convert_to_unified_body),
);
}
let accept_header = req
.headers()
.get(ACCEPT)
.and_then(|accept| accept.to_str().ok())
.map(normalize_header_value)
.unwrap_or_else(|| "application/json".to_string());
let (accept_mode, accepts_sse) = parse_mcp_accept_header(&accept_header);
debug!(
"POST request Accept header: '{}', mode: {:?}, will use SSE for tool calls: {}",
accept_header, accept_mode, accepts_sse
);
let body = req.into_body();
let body_bytes = match body.collect().await {
Ok(collected) => collected.to_bytes(),
Err(err) => {
error!("Failed to read request body: {}", err);
return Ok(bad_request_response("Failed to read request body")
.map(convert_to_unified_body));
}
};
if body_bytes.len() > self.config.max_body_size {
warn!("Request body too large: {} bytes", body_bytes.len());
return Ok(Response::builder()
.status(StatusCode::PAYLOAD_TOO_LARGE)
.header(CONTENT_TYPE, "application/json")
.body(convert_to_unified_body(Full::new(Bytes::from(
"Request body too large",
))))
.unwrap());
}
let body_str = match std::str::from_utf8(&body_bytes) {
Ok(s) => s,
Err(err) => {
error!("Invalid UTF-8 in request body: {}", err);
return Ok(bad_request_response("Request body must be valid UTF-8")
.map(convert_to_unified_body));
}
};
debug!("Received JSON-RPC request: {}", body_str);
let message = match parse_json_rpc_message(body_str) {
Ok(msg) => msg,
Err(rpc_err) => {
error!("JSON-RPC parse error: {}", rpc_err);
let error_response =
serde_json::to_string(&rpc_err).unwrap_or_else(|_| "{}".to_string());
return Ok(Response::builder()
.status(StatusCode::OK) .header(CONTENT_TYPE, "application/json")
.body(convert_to_unified_body(Full::new(Bytes::from(
error_response,
))))
.unwrap());
}
};
let pre_session_extensions = if self.middleware_stack.has_pre_session_middleware() {
let method_name = match &message {
JsonRpcMessage::Request(req) => req.method.as_str(),
JsonRpcMessage::Notification(notif) => notif.method.as_str(),
};
let bearer_token = headers
.get("authorization")
.and_then(|v| extract_bearer_token(v));
let mut pre_ctx = crate::middleware::RequestContext::new(method_name, None);
if let Some(ref token) = bearer_token {
pre_ctx.set_bearer_token(token.clone());
}
for (k, v) in &headers {
if k.eq_ignore_ascii_case("authorization") && is_bearer_scheme(v) {
continue;
}
pre_ctx.add_metadata(k.clone(), serde_json::json!(v));
}
match self
.middleware_stack
.execute_before_session(&mut pre_ctx)
.await
{
Ok(()) => Some(pre_ctx.take_extensions()),
Err(crate::middleware::MiddlewareError::HttpChallenge {
status,
www_authenticate,
body,
}) => {
let body_str = body.unwrap_or_default();
return Ok(Response::builder()
.status(StatusCode::from_u16(status).unwrap_or(StatusCode::UNAUTHORIZED))
.header("WWW-Authenticate", &www_authenticate)
.header("Cache-Control", "no-store")
.header(CONTENT_TYPE, "application/json")
.body(convert_to_unified_body(Full::new(Bytes::from(body_str))))
.unwrap());
}
Err(other_err) => {
if let JsonRpcMessage::Request(ref req) = message {
let response =
Self::map_middleware_error_to_jsonrpc(other_err, req.id.clone());
let response_json =
serde_json::to_string(&response).unwrap_or_else(|_| "{}".to_string());
return Ok(Response::builder()
.status(StatusCode::OK)
.header(CONTENT_TYPE, "application/json")
.body(convert_to_unified_body(Full::new(Bytes::from(
response_json,
))))
.unwrap());
} else {
return Ok(Response::builder()
.status(StatusCode::FORBIDDEN)
.body(convert_to_unified_body(Full::new(Bytes::from(
other_err.to_string(),
))))
.unwrap());
}
}
}
} else {
None
};
let (message_result, response_session_id, method_name, collected_notifications) =
match message {
JsonRpcMessage::Request(request) => {
debug!("Processing JSON-RPC request: method={}", request.method);
let method_name = request.method.clone();
let (response, response_session_id, inline_notifications) =
if request.method == "initialize" {
debug!(
"Handling initialize request - creating new session via session storage"
);
let capabilities = ServerCapabilities::default();
match self.session_storage.create_session(capabilities).await {
Ok(session_info) => {
debug!(
"Created new session via session storage: {}",
session_info.session_id
);
let broadcaster: SharedNotificationBroadcaster =
Arc::new(StreamManagerNotificationBroadcaster::new(Arc::clone(
&self.stream_manager,
)));
let broadcaster_any =
Arc::new(broadcaster) as Arc<dyn std::any::Any + Send + Sync>;
let session_context = SessionContext {
session_id: session_info.session_id.clone(),
metadata: std::collections::HashMap::new(),
broadcaster: Some(broadcaster_any),
timestamp: chrono::Utc::now().timestamp_millis() as u64,
extensions: std::collections::HashMap::new(),
};
let (response, _) = self
.run_middleware_and_dispatch(
request,
headers.clone(),
session_context,
pre_session_extensions.clone(),
)
.await;
(response, Some(session_info.session_id), Vec::new())
}
Err(err) => {
error!("Failed to create session during initialize: {}", err);
let error_msg = format!("Session creation failed: {}", err);
let error_response = turul_mcp_json_rpc_server::JsonRpcMessage::error(
turul_mcp_json_rpc_server::JsonRpcError::internal_error(
Some(request.id),
Some(error_msg),
),
);
(error_response, None, Vec::new())
}
}
} else {
if let Some(ref session_id_str) = session_id {
if let Err(err) = self.validate_session_exists(session_id_str).await {
warn!(
"Session validation failed for session '{}': {}",
session_id_str, err
);
let body = serde_json::json!({
"error": {
"code": 404,
"message": format!("Session validation failed: {}", err)
}
})
.to_string();
return Ok(Response::builder()
.status(StatusCode::NOT_FOUND)
.header(CONTENT_TYPE, "application/json")
.body(convert_to_unified_body(Full::new(Bytes::from(body))))
.unwrap());
}
}
let session_context = if let Some(ref session_id_str) = session_id {
debug!("Processing request with session: {}", session_id_str);
let broadcaster: SharedNotificationBroadcaster =
Arc::new(StreamManagerNotificationBroadcaster::new(Arc::clone(
&self.stream_manager,
)));
let broadcaster_any =
Arc::new(broadcaster) as Arc<dyn std::any::Any + Send + Sync>;
Some(SessionContext {
session_id: session_id_str.clone(),
metadata: std::collections::HashMap::new(),
broadcaster: Some(broadcaster_any),
timestamp: chrono::Utc::now().timestamp_millis() as u64,
extensions: std::collections::HashMap::new(),
})
} else {
debug!("Processing request without session (lenient mode)");
None
};
let is_tool_call = method_name == "tools/call";
let temp_connection = if is_tool_call {
if let Some(ref sid) = session_id {
let (notify_tx, notify_rx) =
tokio::sync::mpsc::channel::<turul_mcp_session_storage::SseEvent>(
64,
);
let conn_id =
format!("post-{}", uuid::Uuid::now_v7().as_simple());
match self
.stream_manager
.register_streaming_connection(sid, conn_id.clone(), notify_tx)
.await
{
Ok(()) => {
debug!(
"Registered temporary POST connection for tool call: session={}, connection={}",
sid, conn_id
);
Some((conn_id, notify_rx, sid.clone()))
}
Err(e) => {
debug!(
"Could not register temporary connection (non-fatal): {}",
e
);
None
}
}
} else {
None
}
} else {
None
};
let (response, _stashed_injection) = if let Some(ctx) = session_context {
self.run_middleware_and_dispatch(
request,
headers.clone(),
ctx,
pre_session_extensions.clone(),
)
.await
} else {
(self.dispatcher.handle_request(request).await, None)
};
let mut inline_notifications = Vec::new();
if let Some((conn_id, mut notify_rx, sid)) = temp_connection {
while let Ok(event) = notify_rx.try_recv() {
if event.event_type != "ping" && event.event_type != "keepalive" {
debug!(
"Captured inline notification: session={}, event_type={}",
sid, event.event_type
);
inline_notifications.push(event);
}
}
self.stream_manager
.unregister_connection(&sid, &conn_id)
.await;
}
(response, session_id, inline_notifications)
};
let message_result = match response {
turul_mcp_json_rpc_server::JsonRpcMessage::Response(resp) => {
JsonRpcMessageResult::Response(resp)
}
turul_mcp_json_rpc_server::JsonRpcMessage::Error(err) => {
JsonRpcMessageResult::Error(err)
}
};
(
message_result,
response_session_id,
Some(method_name),
inline_notifications,
)
}
JsonRpcMessage::Notification(notification) => {
debug!(
"Processing JSON-RPC notification: method={}",
notification.method
);
let method_name = notification.method.clone();
let session_context = if let Some(ref session_id_str) = session_id {
debug!("Processing notification with session: {}", session_id_str);
let broadcaster: SharedNotificationBroadcaster = Arc::new(
StreamManagerNotificationBroadcaster::new(Arc::clone(&self.stream_manager)),
);
let broadcaster_any =
Arc::new(broadcaster) as Arc<dyn std::any::Any + Send + Sync>;
Some(SessionContext {
session_id: session_id_str.clone(),
metadata: std::collections::HashMap::new(),
broadcaster: Some(broadcaster_any),
timestamp: chrono::Utc::now().timestamp_millis() as u64,
extensions: std::collections::HashMap::new(),
})
} else {
debug!("Processing notification without session (lenient mode)");
None
};
let result = self
.dispatcher
.handle_notification_with_context(notification, session_context)
.await;
if let Err(err) = result {
error!("Notification handling error: {}", err);
}
(
JsonRpcMessageResult::NoResponse,
session_id.clone(),
Some(method_name),
Vec::new(),
)
}
};
match message_result {
JsonRpcMessageResult::Response(response) => {
let is_tool_call = method_name.as_ref().is_some_and(|m| m == "tools/call");
debug!(
"Decision point: method={:?}, accept_mode={:?}, accepts_sse={}, server_post_sse_enabled={}, session_id={:?}, is_tool_call={}",
method_name,
accept_mode,
accepts_sse,
self.config.enable_post_sse,
response_session_id,
is_tool_call
);
let should_use_sse = match accept_mode {
AcceptMode::JsonOnly => false, AcceptMode::Invalid => false, AcceptMode::Compliant => {
self.config.enable_post_sse && accepts_sse && is_tool_call
} AcceptMode::SseOnly => self.config.enable_post_sse && accepts_sse, };
if should_use_sse && response_session_id.is_some() {
debug!(
"📡 Creating POST SSE stream (mode: {:?}) for tool call with {} inline notifications",
accept_mode,
collected_notifications.len()
);
let sse_result = if !collected_notifications.is_empty() {
self.stream_manager
.create_post_sse_stream_with_notifications(
response_session_id.clone().unwrap(),
response.clone(),
collected_notifications,
)
.await
} else {
self.stream_manager
.create_post_sse_stream(
response_session_id.clone().unwrap(),
response.clone(),
)
.await
};
match sse_result {
Ok(sse_response) => {
debug!("POST SSE stream created successfully");
Ok(sse_response
.map(|body| body.map_err(|never| match never {}).boxed_unsync()))
}
Err(e) => {
warn!(
"Failed to create POST SSE stream, falling back to JSON: {}",
e
);
Ok(
jsonrpc_response_with_session(response, response_session_id)?
.map(convert_to_unified_body),
)
}
}
} else {
debug!(
"📄 Returning standard JSON response (mode: {:?}) for method: {:?}",
accept_mode, method_name
);
Ok(
jsonrpc_response_with_session(response, response_session_id)?
.map(convert_to_unified_body),
)
}
}
JsonRpcMessageResult::Error(error) => {
warn!("Sending JSON-RPC error response");
let error_json = serde_json::to_string(&error)?;
Ok(Response::builder()
.status(StatusCode::OK) .header(CONTENT_TYPE, "application/json")
.body(convert_to_unified_body(Full::new(Bytes::from(error_json))))
.unwrap())
}
JsonRpcMessageResult::NoResponse => {
Ok(jsonrpc_notification_response()?.map(convert_to_unified_body))
}
}
}
async fn handle_sse_request<B>(&self, req: Request<B>) -> Result<Response<UnifiedMcpBody>>
where
B: http_body::Body<Data = bytes::Bytes, Error = hyper::Error> + Send + 'static,
{
let headers = req.headers();
let accept = headers
.get(ACCEPT)
.and_then(|accept| accept.to_str().ok())
.map(normalize_header_value)
.unwrap_or_default();
if !accept.contains("text/event-stream") {
warn!(
"GET request received without SSE support - header does not contain 'text/event-stream'"
);
let error = JsonRpcError::new(
None,
JsonRpcErrorObject::server_error(
-32001,
"SSE not accepted - missing 'text/event-stream' in Accept header",
None,
),
);
return jsonrpc_error_to_unified_body(error);
}
if !self.config.enable_get_sse {
warn!("GET SSE request received but GET SSE is disabled on server");
let error = JsonRpcError::new(
None,
JsonRpcErrorObject::server_error(
-32003,
"GET SSE is disabled on this server",
None,
),
);
return jsonrpc_error_to_unified_body(error);
}
let protocol_version = extract_protocol_version(headers);
let session_id = extract_session_id(headers);
debug!(
"GET SSE request - Protocol: {}, Session: {:?}",
protocol_version, session_id
);
let session_id = match session_id {
Some(id) => id,
None => {
warn!("Missing Mcp-Session-Id header for SSE request");
let error = JsonRpcError::new(
None,
JsonRpcErrorObject::server_error(-32002, "Missing Mcp-Session-Id header", None),
);
return jsonrpc_error_to_unified_body(error);
}
};
if let Err(err) = self.validate_session_exists(&session_id).await {
warn!(
"Session validation failed for session '{}': {}",
session_id, err
);
let body = serde_json::json!({
"error": {
"code": 404,
"message": format!("Session validation failed: {}", err)
}
})
.to_string();
return Ok(Response::builder()
.status(StatusCode::NOT_FOUND)
.header("content-type", "application/json")
.body(convert_to_unified_body(Full::new(Bytes::from(body))))
.unwrap());
}
let last_event_id = extract_last_event_id(headers);
let connection_id = Uuid::now_v7().as_simple().to_string();
debug!(
"Creating SSE stream for session: {} with connection: {}, last_event_id: {:?}",
session_id, connection_id, last_event_id
);
match self
.stream_manager
.handle_sse_connection(session_id, connection_id, last_event_id)
.await
{
Ok(response) => Ok(response),
Err(err) => {
error!("Failed to create SSE connection: {}", err);
let error = JsonRpcError::new(
None,
JsonRpcErrorObject::internal_error(Some(format!(
"SSE connection failed: {}",
err
))),
);
jsonrpc_error_to_unified_body(error)
}
}
}
async fn handle_delete_request<B>(&self, req: Request<B>) -> Result<Response<JsonRpcBody>>
where
B: http_body::Body<Data = bytes::Bytes, Error = hyper::Error> + Send + 'static,
{
let session_id = extract_session_id(req.headers());
debug!("DELETE request - Session: {:?}", session_id);
if let Some(session_id) = session_id {
let closed_connections = self
.stream_manager
.close_session_connections(&session_id)
.await;
debug!(
"Closed {} SSE connections for session: {}",
closed_connections, session_id
);
match self.session_storage.get_session(&session_id).await {
Ok(Some(mut session_info)) => {
session_info
.state
.insert("terminated".to_string(), serde_json::Value::Bool(true));
session_info.state.insert(
"terminated_at".to_string(),
serde_json::Value::Number(serde_json::Number::from(
chrono::Utc::now().timestamp_millis(),
)),
);
session_info.touch();
match self.session_storage.update_session(session_info).await {
Ok(()) => {
debug!(
"Session {} marked as terminated (TTL will handle cleanup)",
session_id
);
Ok(Response::builder()
.status(StatusCode::OK)
.body(Full::new(Bytes::from("Session terminated")))
.unwrap())
}
Err(err) => {
error!(
"Error marking session {} as terminated: {}",
session_id, err
);
match self.session_storage.delete_session(&session_id).await {
Ok(_) => {
debug!("Session {} deleted as fallback", session_id);
Ok(Response::builder()
.status(StatusCode::OK)
.body(Full::new(Bytes::from("Session removed")))
.unwrap())
}
Err(delete_err) => {
error!(
"Error deleting session {} as fallback: {}",
session_id, delete_err
);
Ok(Response::builder()
.status(StatusCode::INTERNAL_SERVER_ERROR)
.body(Full::new(Bytes::from("Session termination error")))
.unwrap())
}
}
}
}
}
Ok(None) => Ok(Response::builder()
.status(StatusCode::NOT_FOUND)
.body(Full::new(Bytes::from("Session not found")))
.unwrap()),
Err(err) => {
error!(
"Error retrieving session {} for termination: {}",
session_id, err
);
Ok(Response::builder()
.status(StatusCode::INTERNAL_SERVER_ERROR)
.body(Full::new(Bytes::from("Session lookup error")))
.unwrap())
}
}
} else {
Ok(Response::builder()
.status(StatusCode::BAD_REQUEST)
.body(Full::new(Bytes::from("Missing Mcp-Session-Id header")))
.unwrap())
}
}
fn handle_preflight(&self) -> Response<JsonRpcBody> {
options_response()
}
fn method_not_allowed(&self) -> Response<JsonRpcBody> {
method_not_allowed_response()
}
async fn validate_session_exists(&self, session_id: &str) -> Result<()> {
match self.session_storage.get_session(session_id).await {
Ok(Some(session_info)) => {
if session_info.is_terminated() {
error!("Session '{}' has been terminated", session_id);
return Err(crate::HttpMcpError::InvalidRequest(format!(
"Session '{}' has been terminated. Create a new session to continue.",
session_id
)));
}
if let Some(ref current_fp) = self.tool_fingerprint {
if let Some(stored_fp) = session_info.state.get("mcp:tool_fingerprint") {
if stored_fp.as_str() != Some(current_fp.as_str()) {
info!(
"Tool fingerprint updated for session '{}' (tools changed since session created)",
session_id
);
let _ = self.session_storage
.set_session_state(session_id, "mcp:tool_fingerprint", serde_json::json!(current_fp))
.await;
if let Some(ref notifier) = self.tool_notifier {
notifier.notify_tools_changed(session_id).await
.map_err(|e| crate::HttpMcpError::Mcp(
McpError::tool_execution(&format!(
"Notification persistence failed for session {}: {}", session_id, e
))
))?;
}
}
} else {
info!("Session '{}' has no tool fingerprint, storing current", session_id);
let _ = self.session_storage
.set_session_state(session_id, "mcp:tool_fingerprint", serde_json::json!(current_fp))
.await;
}
}
debug!("Session validation successful: {}", session_id);
Ok(())
}
Ok(None) => {
error!("Session not found: {}", session_id);
Err(crate::HttpMcpError::InvalidRequest(format!(
"Session '{}' not found. Sessions must be created via initialize request first.",
session_id
)))
}
Err(err) => {
error!("Failed to validate session {}: {}", session_id, err);
Err(crate::HttpMcpError::InvalidRequest(format!(
"Session validation failed: {}",
err
)))
}
}
}
async fn run_middleware_and_dispatch(
&self,
request: turul_mcp_json_rpc_server::JsonRpcRequest,
headers: HashMap<String, String>,
session: turul_mcp_json_rpc_server::SessionContext,
pre_session_extensions: Option<HashMap<String, serde_json::Value>>,
) -> (
turul_mcp_json_rpc_server::JsonRpcMessage,
Option<crate::middleware::SessionInjection>,
) {
if self.middleware_stack.is_empty() {
let result = self
.dispatcher
.handle_request_with_context(request, session)
.await;
return (result, None);
}
let normalized_headers: HashMap<String, String> = headers
.iter()
.map(|(k, v)| (k.to_lowercase(), v.clone()))
.collect();
let method = request.method.clone();
let session_id = session.session_id.clone();
let params = request.params.clone().map(|p| match p {
turul_mcp_json_rpc_server::RequestParams::Object(map) => {
serde_json::Value::Object(map.into_iter().collect())
}
turul_mcp_json_rpc_server::RequestParams::Array(arr) => serde_json::Value::Array(arr),
});
let mut ctx = crate::middleware::RequestContext::new(&method, params);
if let Some(ext) = pre_session_extensions {
for (k, v) in ext {
ctx.set_extension(k, v);
}
}
for (k, v) in normalized_headers {
if k == "authorization" && is_bearer_scheme(&v) {
continue;
}
ctx.add_metadata(k, serde_json::json!(v));
}
let session_view = crate::middleware::StorageBackedSessionView::new(
session_id.clone(),
Arc::clone(&self.session_storage),
);
let injection = match self
.middleware_stack
.execute_before(&mut ctx, Some(&session_view))
.await
{
Ok(inj) => inj,
Err(err) => {
return (Self::map_middleware_error_to_jsonrpc(err, request.id), None);
}
};
if !injection.is_empty() {
for (key, value) in injection.state() {
if let Err(e) = session_view.set_state(key, value.clone()).await {
tracing::warn!("Failed to apply injection state '{}': {}", key, e);
}
}
for (key, value) in injection.metadata() {
if let Err(e) = session_view.set_metadata(key, value.clone()).await {
tracing::warn!("Failed to apply injection metadata '{}': {}", key, e);
}
}
}
let mut session = session;
session.extensions = ctx.extensions().clone();
let request_id = request.id.clone();
let result = self
.dispatcher
.handle_request_with_context(request, session)
.await;
let mut dispatcher_result = match &result {
turul_mcp_json_rpc_server::JsonRpcMessage::Response(resp) => match &resp.result {
turul_mcp_json_rpc_server::response::ResponseResult::Success(val) => {
crate::middleware::DispatcherResult::Success(val.clone())
}
turul_mcp_json_rpc_server::response::ResponseResult::Null => {
crate::middleware::DispatcherResult::Success(serde_json::Value::Null)
}
},
turul_mcp_json_rpc_server::JsonRpcMessage::Error(err) => {
crate::middleware::DispatcherResult::Error(err.error.message.clone())
}
};
match self
.middleware_stack
.execute_after(&ctx, &mut dispatcher_result)
.await
{
Ok(()) => {
let result = Self::apply_dispatcher_result(result, dispatcher_result);
(result, None)
}
Err(middleware_err) => (
Self::map_middleware_error_to_jsonrpc(middleware_err, request_id),
None,
),
}
}
fn apply_dispatcher_result(
result: turul_mcp_json_rpc_server::JsonRpcMessage,
dispatcher_result: crate::middleware::DispatcherResult,
) -> turul_mcp_json_rpc_server::JsonRpcMessage {
match dispatcher_result {
crate::middleware::DispatcherResult::Success(val) => match result {
turul_mcp_json_rpc_server::JsonRpcMessage::Response(mut resp) => {
resp.result = turul_mcp_json_rpc_server::response::ResponseResult::Success(val);
turul_mcp_json_rpc_server::JsonRpcMessage::Response(resp)
}
turul_mcp_json_rpc_server::JsonRpcMessage::Error(err) => {
match err.id {
Some(id) => turul_mcp_json_rpc_server::JsonRpcMessage::Response(
turul_mcp_json_rpc_server::response::JsonRpcResponse::success(id, val),
),
None => turul_mcp_json_rpc_server::JsonRpcMessage::Error(err),
}
}
},
crate::middleware::DispatcherResult::Error(msg) => match result {
turul_mcp_json_rpc_server::JsonRpcMessage::Response(resp) => {
turul_mcp_json_rpc_server::JsonRpcMessage::Error(
turul_mcp_json_rpc_server::error::JsonRpcError::new(
Some(resp.id),
turul_mcp_json_rpc_server::error::JsonRpcErrorObject::internal_error(
Some(msg),
),
),
)
}
turul_mcp_json_rpc_server::JsonRpcMessage::Error(mut err) => {
err.error.message = msg;
turul_mcp_json_rpc_server::JsonRpcMessage::Error(err)
}
},
}
}
fn map_middleware_error_to_jsonrpc(
err: crate::middleware::MiddlewareError,
request_id: turul_mcp_json_rpc_server::RequestId,
) -> turul_mcp_json_rpc_server::JsonRpcMessage {
use crate::middleware::MiddlewareError;
use crate::middleware::error::error_codes;
let (code, message, data) = match err {
MiddlewareError::Unauthenticated(msg) => (error_codes::UNAUTHENTICATED, msg, None),
MiddlewareError::Unauthorized(msg) => (error_codes::UNAUTHORIZED, msg, None),
MiddlewareError::RateLimitExceeded {
message,
retry_after,
} => {
let data = retry_after.map(|s| serde_json::json!({"retryAfter": s}));
(error_codes::RATE_LIMIT_EXCEEDED, message, data)
}
MiddlewareError::InvalidRequest(msg) => (error_codes::INVALID_REQUEST, msg, None),
MiddlewareError::Internal(msg) => (error_codes::INTERNAL_ERROR, msg, None),
MiddlewareError::Custom { message, .. } => (error_codes::INTERNAL_ERROR, message, None),
MiddlewareError::HttpChallenge { .. } => {
unreachable!(
"HttpChallenge must be caught at transport level before JSON-RPC dispatch"
)
}
};
let error_obj = if let Some(d) = data {
turul_mcp_json_rpc_server::error::JsonRpcErrorObject::server_error(
code,
&message,
Some(d),
)
} else {
turul_mcp_json_rpc_server::error::JsonRpcErrorObject::server_error(
code,
&message,
None::<serde_json::Value>,
)
};
turul_mcp_json_rpc_server::JsonRpcMessage::Error(
turul_mcp_json_rpc_server::JsonRpcError::new(Some(request_id), error_obj),
)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_validate_session_exists_rejects_terminated() {
let storage: Arc<turul_mcp_session_storage::BoxedSessionStorage> =
Arc::new(InMemorySessionStorage::new());
let mut session = storage
.create_session(turul_mcp_protocol::ServerCapabilities::default())
.await
.unwrap();
let session_id = session.session_id.clone();
session
.state
.insert("terminated".to_string(), serde_json::json!(true));
storage.update_session(session).await.unwrap();
let dispatcher = Arc::new(JsonRpcDispatcher::<McpError>::default());
let handler = SessionMcpHandler::with_storage(
crate::server::ServerConfig::default(),
dispatcher,
storage,
crate::stream_manager::StreamConfig::default(),
Arc::new(crate::middleware::MiddlewareStack::new()),
);
let result = handler.validate_session_exists(&session_id).await;
assert!(result.is_err(), "Expected Err for terminated session");
let err_msg = format!("{}", result.unwrap_err());
assert!(
err_msg.to_lowercase().contains("terminated"),
"Error must mention 'terminated', got: {}",
err_msg
);
}
#[tokio::test]
async fn test_fingerprint_mismatch_updates_and_continues() {
let storage: Arc<turul_mcp_session_storage::BoxedSessionStorage> =
Arc::new(InMemorySessionStorage::new());
let session = storage
.create_session(turul_mcp_protocol::ServerCapabilities::default())
.await
.unwrap();
let session_id = session.session_id.clone();
storage
.set_session_state(&session_id, "mcp:tool_fingerprint", serde_json::json!("fp_v1"))
.await
.unwrap();
let dispatcher = Arc::new(JsonRpcDispatcher::<McpError>::default());
let handler = SessionMcpHandler::with_storage(
crate::server::ServerConfig::default(),
dispatcher,
storage.clone(),
crate::stream_manager::StreamConfig::default(),
Arc::new(crate::middleware::MiddlewareStack::new()),
)
.with_tool_fingerprint(Some("fp_v2".to_string()));
let result = handler.validate_session_exists(&session_id).await;
assert!(result.is_ok(), "Fingerprint mismatch must NOT reject session");
let stored_fp = storage
.get_session_state(&session_id, "mcp:tool_fingerprint")
.await
.unwrap();
assert_eq!(
stored_fp.and_then(|v| v.as_str().map(String::from)),
Some("fp_v2".to_string()),
"Stored fingerprint must be updated to current"
);
}
#[tokio::test]
async fn test_validate_session_accepts_matching_fingerprint() {
let storage: Arc<turul_mcp_session_storage::BoxedSessionStorage> =
Arc::new(InMemorySessionStorage::new());
let session = storage
.create_session(turul_mcp_protocol::ServerCapabilities::default())
.await
.unwrap();
let session_id = session.session_id.clone();
storage
.set_session_state(&session_id, "mcp:tool_fingerprint", serde_json::json!("fp_v1"))
.await
.unwrap();
let dispatcher = Arc::new(JsonRpcDispatcher::<McpError>::default());
let handler = SessionMcpHandler::with_storage(
crate::server::ServerConfig::default(),
dispatcher,
storage,
crate::stream_manager::StreamConfig::default(),
Arc::new(crate::middleware::MiddlewareStack::new()),
)
.with_tool_fingerprint(Some("fp_v1".to_string()));
let result = handler.validate_session_exists(&session_id).await;
assert!(result.is_ok(), "Matching fingerprint should accept session");
}
#[tokio::test]
async fn test_legacy_session_without_fingerprint_gets_current_stored() {
let storage: Arc<turul_mcp_session_storage::BoxedSessionStorage> =
Arc::new(InMemorySessionStorage::new());
let session = storage
.create_session(turul_mcp_protocol::ServerCapabilities::default())
.await
.unwrap();
let session_id = session.session_id.clone();
let dispatcher = Arc::new(JsonRpcDispatcher::<McpError>::default());
let handler = SessionMcpHandler::with_storage(
crate::server::ServerConfig::default(),
dispatcher,
storage.clone(),
crate::stream_manager::StreamConfig::default(),
Arc::new(crate::middleware::MiddlewareStack::new()),
)
.with_tool_fingerprint(Some("current_fp".to_string()));
let result = handler.validate_session_exists(&session_id).await;
assert!(result.is_ok(), "Legacy session must NOT be rejected");
let stored_fp = storage
.get_session_state(&session_id, "mcp:tool_fingerprint")
.await
.unwrap();
assert_eq!(
stored_fp.and_then(|v| v.as_str().map(String::from)),
Some("current_fp".to_string()),
"Current fingerprint must be stored for legacy session"
);
}
#[tokio::test]
async fn test_validate_session_skips_fingerprint_when_handler_has_none() {
let storage: Arc<turul_mcp_session_storage::BoxedSessionStorage> =
Arc::new(InMemorySessionStorage::new());
let session = storage
.create_session(turul_mcp_protocol::ServerCapabilities::default())
.await
.unwrap();
let session_id = session.session_id.clone();
let dispatcher = Arc::new(JsonRpcDispatcher::<McpError>::default());
let handler = SessionMcpHandler::with_storage(
crate::server::ServerConfig::default(),
dispatcher,
storage,
crate::stream_manager::StreamConfig::default(),
Arc::new(crate::middleware::MiddlewareStack::new()),
);
let result = handler.validate_session_exists(&session_id).await;
assert!(result.is_ok(), "Handler without fingerprint should accept all sessions");
}
#[tokio::test]
async fn test_persisted_session_fingerprint_updated_on_mismatch() {
let storage: Arc<turul_mcp_session_storage::BoxedSessionStorage> =
Arc::new(InMemorySessionStorage::new());
let session = storage
.create_session(turul_mcp_protocol::ServerCapabilities::default())
.await
.unwrap();
let session_id = session.session_id.clone();
storage
.set_session_state(&session_id, "mcp:tool_fingerprint", serde_json::json!("deploy_v1"))
.await
.unwrap();
let dispatcher = Arc::new(JsonRpcDispatcher::<McpError>::default());
let handler = SessionMcpHandler::with_storage(
crate::server::ServerConfig::default(),
dispatcher,
storage.clone(),
crate::stream_manager::StreamConfig::default(),
Arc::new(crate::middleware::MiddlewareStack::new()),
)
.with_tool_fingerprint(Some("deploy_v2".to_string()));
let result = handler.validate_session_exists(&session_id).await;
assert!(result.is_ok(), "Session must continue on fingerprint mismatch");
let stored_fp = storage
.get_session_state(&session_id, "mcp:tool_fingerprint")
.await
.unwrap();
assert_eq!(
stored_fp.and_then(|v| v.as_str().map(String::from)),
Some("deploy_v2".to_string()),
"Stored fingerprint must be updated to server's current"
);
let still_stored = storage.get_session(&session_id).await.unwrap();
assert!(still_stored.is_some(), "Session must remain in storage");
}
#[tokio::test]
async fn test_fingerprint_update_is_idempotent() {
let storage: Arc<turul_mcp_session_storage::BoxedSessionStorage> =
Arc::new(InMemorySessionStorage::new());
let session = storage
.create_session(turul_mcp_protocol::ServerCapabilities::default())
.await
.unwrap();
let session_id = session.session_id.clone();
storage
.set_session_state(&session_id, "mcp:tool_fingerprint", serde_json::json!("old_fp"))
.await
.unwrap();
let dispatcher = Arc::new(JsonRpcDispatcher::<McpError>::default());
let handler = SessionMcpHandler::with_storage(
crate::server::ServerConfig::default(),
dispatcher,
storage.clone(),
crate::stream_manager::StreamConfig::default(),
Arc::new(crate::middleware::MiddlewareStack::new()),
)
.with_tool_fingerprint(Some("new_fp".to_string()));
let result = handler.validate_session_exists(&session_id).await;
assert!(result.is_ok(), "First request with mismatch should succeed");
let result2 = handler.validate_session_exists(&session_id).await;
assert!(result2.is_ok(), "Second request should proceed (fingerprints match)");
let stored_fp = storage
.get_session_state(&session_id, "mcp:tool_fingerprint")
.await
.unwrap();
assert_eq!(
stored_fp.and_then(|v| v.as_str().map(String::from)),
Some("new_fp".to_string()),
);
}
}