use std::sync::Arc;
use lambda_http::{Body as LambdaBody, Request as LambdaRequest, Response as LambdaResponse};
use tracing::{debug, info};
use turul_http_mcp_server::{
ServerConfig, SessionMcpHandler, StreamConfig, StreamManager, StreamableHttpHandler,
};
use turul_mcp_json_rpc_server::JsonRpcDispatcher;
use turul_mcp_protocol::{McpError, ServerCapabilities};
use turul_mcp_session_storage::BoxedSessionStorage;
use crate::error::Result;
#[cfg(feature = "cors")]
use crate::cors::{CorsConfig, create_preflight_response, inject_cors_headers};
#[derive(Clone)]
pub struct LambdaMcpHandler {
session_handler: SessionMcpHandler,
streamable_handler: StreamableHttpHandler,
#[allow(dead_code)]
sse_enabled: bool,
route_registry: Arc<turul_http_mcp_server::RouteRegistry>,
#[cfg(feature = "dynamic-tools")]
tool_registry: Option<Arc<turul_mcp_server::ToolRegistry>>,
#[cfg(feature = "cors")]
cors_config: Option<CorsConfig>,
}
impl LambdaMcpHandler {
#[allow(clippy::too_many_arguments)]
pub fn new(
dispatcher: JsonRpcDispatcher<McpError>,
session_storage: Arc<BoxedSessionStorage>,
stream_manager: Arc<StreamManager>,
config: ServerConfig,
stream_config: StreamConfig,
_implementation: turul_mcp_protocol::Implementation,
capabilities: ServerCapabilities,
sse_enabled: bool,
#[cfg(feature = "cors")] cors_config: Option<CorsConfig>,
) -> Self {
let dispatcher = Arc::new(dispatcher);
let middleware_stack = Arc::new(turul_http_mcp_server::middleware::MiddlewareStack::new());
let session_handler = SessionMcpHandler::with_shared_stream_manager(
config.clone(),
dispatcher.clone(),
session_storage.clone(),
stream_config.clone(),
stream_manager.clone(),
middleware_stack.clone(),
);
let streamable_handler = StreamableHttpHandler::new(
Arc::new(config.clone()),
dispatcher.clone(),
session_storage.clone(),
stream_manager.clone(),
capabilities.clone(),
middleware_stack,
None, );
Self {
session_handler,
streamable_handler,
sse_enabled,
route_registry: Arc::new(turul_http_mcp_server::RouteRegistry::new()),
#[cfg(feature = "dynamic-tools")]
tool_registry: None,
#[cfg(feature = "cors")]
cors_config,
}
}
#[allow(clippy::too_many_arguments)]
pub fn with_shared_stream_manager(
config: ServerConfig,
dispatcher: Arc<JsonRpcDispatcher<McpError>>,
session_storage: Arc<BoxedSessionStorage>,
stream_manager: Arc<StreamManager>,
stream_config: StreamConfig,
_implementation: turul_mcp_protocol::Implementation,
capabilities: ServerCapabilities,
sse_enabled: bool,
) -> Self {
let middleware_stack = Arc::new(turul_http_mcp_server::middleware::MiddlewareStack::new());
let session_handler = SessionMcpHandler::with_shared_stream_manager(
config.clone(),
dispatcher.clone(),
session_storage.clone(),
stream_config.clone(),
stream_manager.clone(),
middleware_stack.clone(),
);
let streamable_handler = StreamableHttpHandler::new(
Arc::new(config),
dispatcher,
session_storage,
stream_manager,
capabilities,
middleware_stack,
None, );
Self {
session_handler,
streamable_handler,
sse_enabled,
route_registry: Arc::new(turul_http_mcp_server::RouteRegistry::new()),
#[cfg(feature = "dynamic-tools")]
tool_registry: None,
#[cfg(feature = "cors")]
cors_config: None,
}
}
#[allow(clippy::too_many_arguments)]
pub fn with_middleware(
config: ServerConfig,
dispatcher: Arc<JsonRpcDispatcher<McpError>>,
session_storage: Arc<BoxedSessionStorage>,
stream_manager: Arc<StreamManager>,
stream_config: StreamConfig,
capabilities: ServerCapabilities,
middleware_stack: Arc<turul_http_mcp_server::middleware::MiddlewareStack>,
sse_enabled: bool,
route_registry: Arc<turul_http_mcp_server::RouteRegistry>,
) -> Self {
Self::with_middleware_and_fingerprint(
config,
dispatcher,
session_storage,
stream_manager,
stream_config,
capabilities,
middleware_stack,
sse_enabled,
route_registry,
None,
)
}
#[allow(clippy::too_many_arguments)]
pub fn with_middleware_and_fingerprint(
config: ServerConfig,
dispatcher: Arc<JsonRpcDispatcher<McpError>>,
session_storage: Arc<BoxedSessionStorage>,
stream_manager: Arc<StreamManager>,
stream_config: StreamConfig,
capabilities: ServerCapabilities,
middleware_stack: Arc<turul_http_mcp_server::middleware::MiddlewareStack>,
sse_enabled: bool,
route_registry: Arc<turul_http_mcp_server::RouteRegistry>,
tool_fingerprint: Option<String>,
) -> Self {
let session_handler = SessionMcpHandler::with_shared_stream_manager(
config.clone(),
dispatcher.clone(),
session_storage.clone(),
stream_config.clone(),
stream_manager.clone(),
middleware_stack.clone(),
)
.with_tool_fingerprint(tool_fingerprint.clone());
let streamable_handler = StreamableHttpHandler::new(
Arc::new(config),
dispatcher,
session_storage,
stream_manager,
capabilities,
middleware_stack,
tool_fingerprint,
);
Self {
session_handler,
streamable_handler,
sse_enabled,
route_registry,
#[cfg(feature = "dynamic-tools")]
tool_registry: None,
#[cfg(feature = "cors")]
cors_config: None,
}
}
pub fn with_tool_notifier(
mut self,
notifier: Arc<dyn turul_http_mcp_server::ToolChangeNotifier>,
) -> Self {
self.session_handler = self
.session_handler
.with_tool_notifier(Arc::clone(¬ifier));
self.streamable_handler = self.streamable_handler.with_tool_notifier(notifier);
self
}
#[cfg(feature = "dynamic-tools")]
pub fn with_tool_registry(mut self, registry: Arc<turul_mcp_server::ToolRegistry>) -> Self {
self.tool_registry = Some(registry);
self
}
#[cfg(feature = "cors")]
pub fn with_cors(mut self, cors_config: CorsConfig) -> Self {
self.cors_config = Some(cors_config);
self
}
pub fn get_stream_manager(&self) -> &Arc<StreamManager> {
self.session_handler.get_stream_manager()
}
pub async fn handle(&self, req: LambdaRequest) -> Result<LambdaResponse<LambdaBody>> {
let method = req.method().clone();
let uri = req.uri().clone();
let request_origin = req
.headers()
.get("origin")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
info!(
"🌐 Lambda MCP request: {} {} (origin: {:?})",
method, uri, request_origin
);
#[cfg(feature = "cors")]
if method == http::Method::OPTIONS
&& let Some(ref cors_config) = self.cors_config
{
debug!("Handling CORS preflight request");
return create_preflight_response(cors_config, request_origin.as_deref());
}
#[cfg(feature = "dynamic-tools")]
if let Some(ref registry) = self.tool_registry
&& let Err(e) = registry.check_for_changes().await
{
tracing::warn!(error = %e, "Failed to check for tool changes");
}
let hyper_req = crate::adapter::lambda_to_hyper_request(req)?;
let path = hyper_req.uri().path().to_string();
if !self.route_registry.is_empty() {
match self.route_registry.match_route(&path) {
Ok(Some(route_handler)) => {
debug!("Custom route matched: {}", path);
use http_body_util::BodyExt;
let (parts, body) = hyper_req.into_parts();
let boxed_req = hyper::Request::from_parts(parts, body.boxed_unsync());
let route_resp = route_handler.handle(boxed_req).await;
let mut lambda_resp =
crate::adapter::hyper_to_lambda_response(route_resp).await?;
#[cfg(feature = "cors")]
if let Some(ref cors_config) = self.cors_config {
inject_cors_headers(
&mut lambda_resp,
cors_config,
request_origin.as_deref(),
)?;
}
return Ok(lambda_resp);
}
Ok(None) => {} Err(e) => {
debug!("Route validation error: {}", e);
let route_resp = e.into_response();
let mut lambda_resp =
crate::adapter::hyper_to_lambda_response(route_resp).await?;
#[cfg(feature = "cors")]
if let Some(ref cors_config) = self.cors_config {
inject_cors_headers(
&mut lambda_resp,
cors_config,
request_origin.as_deref(),
)?;
}
return Ok(lambda_resp);
}
}
}
let hyper_resp = self
.session_handler
.handle_mcp_request(hyper_req)
.await
.map_err(|e| crate::error::LambdaError::McpFramework(e.to_string()))?;
let mut lambda_resp = crate::adapter::hyper_to_lambda_response(hyper_resp).await?;
#[cfg(feature = "cors")]
if let Some(ref cors_config) = self.cors_config {
inject_cors_headers(&mut lambda_resp, cors_config, request_origin.as_deref())?;
}
Ok(lambda_resp)
}
pub async fn handle_streaming(
&self,
req: LambdaRequest,
) -> std::result::Result<
lambda_http::Response<
http_body_util::combinators::UnsyncBoxBody<bytes::Bytes, hyper::Error>,
>,
Box<dyn std::error::Error + Send + Sync>,
> {
let method = req.method().clone();
let uri = req.uri().clone();
let request_origin = req
.headers()
.get("origin")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
debug!(
"🌊 Lambda streaming MCP request: {} {} (origin: {:?})",
method, uri, request_origin
);
#[cfg(feature = "cors")]
if method == http::Method::OPTIONS
&& let Some(ref cors_config) = self.cors_config
{
debug!("Handling CORS preflight request (streaming)");
let preflight_response =
create_preflight_response(cors_config, request_origin.as_deref())
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
return Ok(self.convert_lambda_response_to_streaming(preflight_response));
}
#[cfg(feature = "dynamic-tools")]
if let Some(ref registry) = self.tool_registry
&& let Err(e) = registry.check_for_changes().await
{
tracing::warn!(error = %e, "Failed to check for tool changes (streaming)");
}
let hyper_req = crate::adapter::lambda_to_hyper_request(req)
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
let path = hyper_req.uri().path().to_string();
if !self.route_registry.is_empty() {
match self.route_registry.match_route(&path) {
Ok(Some(route_handler)) => {
debug!("Custom route matched (streaming): {}", path);
use http_body_util::BodyExt;
let (parts, body) = hyper_req.into_parts();
let boxed_req = hyper::Request::from_parts(parts, body.boxed_unsync());
let mut route_resp = route_handler.handle(boxed_req).await;
#[cfg(feature = "cors")]
if let Some(ref cors_config) = self.cors_config {
inject_cors_headers(
&mut route_resp,
cors_config,
request_origin.as_deref(),
)
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
}
return Ok(route_resp);
}
Ok(None) => {} Err(e) => {
debug!("Route validation error (streaming): {}", e);
let mut err_resp = e.into_response();
#[cfg(feature = "cors")]
if let Some(ref cors_config) = self.cors_config {
inject_cors_headers(
&mut err_resp,
cors_config,
request_origin.as_deref(),
)
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
}
return Ok(err_resp);
}
}
}
use turul_http_mcp_server::protocol::McpProtocolVersion;
let protocol_version = hyper_req
.headers()
.get("MCP-Protocol-Version")
.and_then(|h| h.to_str().ok())
.and_then(McpProtocolVersion::parse_version)
.unwrap_or(McpProtocolVersion::V2025_06_18);
let hyper_resp = if protocol_version.supports_streamable_http() {
debug!(
"Using StreamableHttpHandler for protocol {}",
protocol_version.to_string()
);
self.streamable_handler.handle_request(hyper_req).await
} else {
debug!(
"Using SessionMcpHandler for legacy protocol {}",
protocol_version.to_string()
);
self.session_handler
.handle_mcp_request(hyper_req)
.await
.map_err(|e| {
Box::new(crate::error::LambdaError::McpFramework(e.to_string()))
as Box<dyn std::error::Error + Send + Sync>
})?
};
let mut lambda_resp = crate::adapter::hyper_to_lambda_streaming(hyper_resp);
#[cfg(feature = "cors")]
if let Some(ref cors_config) = self.cors_config {
inject_cors_headers(&mut lambda_resp, cors_config, request_origin.as_deref())
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
}
Ok(lambda_resp)
}
fn convert_lambda_response_to_streaming(
&self,
lambda_response: LambdaResponse<LambdaBody>,
) -> lambda_http::Response<http_body_util::combinators::UnsyncBoxBody<bytes::Bytes, hyper::Error>>
{
use bytes::Bytes;
use http_body_util::{BodyExt, Full};
let (parts, body) = lambda_response.into_parts();
let body_bytes = match body {
LambdaBody::Empty => Bytes::new(),
LambdaBody::Text(text) => Bytes::from(text),
LambdaBody::Binary(bytes) => Bytes::from(bytes),
_ => Bytes::new(),
};
let streaming_body = Full::new(body_bytes)
.map_err(|e: std::convert::Infallible| match e {})
.boxed_unsync();
lambda_http::Response::from_parts(parts, streaming_body)
}
}
#[cfg(test)]
mod tests {
use super::*;
use http::Request;
use turul_mcp_session_storage::InMemorySessionStorage;
#[tokio::test]
async fn test_handler_creation() {
let session_storage = Arc::new(InMemorySessionStorage::new());
let stream_manager = Arc::new(StreamManager::new(session_storage.clone()));
let dispatcher = JsonRpcDispatcher::new();
let config = ServerConfig::default();
let implementation = turul_mcp_protocol::Implementation::new("test", "1.0.0");
let capabilities = ServerCapabilities::default();
let handler = LambdaMcpHandler::new(
dispatcher,
session_storage,
stream_manager,
config,
StreamConfig::default(),
implementation,
capabilities,
false, #[cfg(feature = "cors")]
None,
);
assert!(!handler.sse_enabled);
}
#[tokio::test]
async fn test_sse_enabled_with_handle_works() {
let session_storage = Arc::new(InMemorySessionStorage::new());
let stream_manager = Arc::new(StreamManager::new(session_storage.clone()));
let dispatcher = JsonRpcDispatcher::new();
let config = ServerConfig::default();
let implementation = turul_mcp_protocol::Implementation::new("test", "1.0.0");
let capabilities = ServerCapabilities::default();
let handler = LambdaMcpHandler::new(
dispatcher,
session_storage,
stream_manager,
config,
StreamConfig::default(),
implementation,
capabilities,
true, #[cfg(feature = "cors")]
None,
);
let lambda_req = Request::builder()
.method("POST")
.uri("/mcp")
.body(LambdaBody::Text(
r#"{"jsonrpc":"2.0","method":"initialize","id":1}"#.to_string(),
))
.unwrap();
let result = handler.handle(lambda_req).await;
assert!(
result.is_ok(),
"handle() should work with SSE enabled for snapshot-based responses"
);
}
#[tokio::test]
async fn test_stream_config_preservation() {
let session_storage = Arc::new(InMemorySessionStorage::new());
let dispatcher = JsonRpcDispatcher::new();
let config = ServerConfig::default();
let implementation = turul_mcp_protocol::Implementation::new("test", "1.0.0");
let capabilities = ServerCapabilities::default();
let custom_stream_config = StreamConfig {
channel_buffer_size: 1024, max_replay_events: 200, keepalive_interval_seconds: 10, cors_origin: "https://custom-test.example.com".to_string(), };
let stream_manager = Arc::new(StreamManager::with_config(
session_storage.clone(),
custom_stream_config.clone(),
));
let handler = LambdaMcpHandler::new(
dispatcher,
session_storage,
stream_manager,
config,
custom_stream_config.clone(),
implementation,
capabilities,
false, #[cfg(feature = "cors")]
None,
);
assert!(!handler.sse_enabled);
let stream_manager = handler.get_stream_manager();
let actual_config = stream_manager.get_config();
assert_eq!(
actual_config.channel_buffer_size, custom_stream_config.channel_buffer_size,
"Custom channel_buffer_size was not propagated correctly"
);
assert_eq!(
actual_config.max_replay_events, custom_stream_config.max_replay_events,
"Custom max_replay_events was not propagated correctly"
);
assert_eq!(
actual_config.keepalive_interval_seconds,
custom_stream_config.keepalive_interval_seconds,
"Custom keepalive_interval_seconds was not propagated correctly"
);
assert_eq!(
actual_config.cors_origin, custom_stream_config.cors_origin,
"Custom cors_origin was not propagated correctly"
);
assert!(Arc::strong_count(stream_manager) >= 1);
}
#[tokio::test]
async fn test_full_builder_chain_stream_config() {
use crate::LambdaMcpServerBuilder;
use turul_mcp_session_storage::InMemorySessionStorage;
let custom_stream_config = turul_http_mcp_server::StreamConfig {
channel_buffer_size: 2048, max_replay_events: 500, keepalive_interval_seconds: 15, cors_origin: "https://full-chain-test.example.com".to_string(),
};
let server = LambdaMcpServerBuilder::new()
.name("full-chain-test")
.version("1.0.0")
.storage(Arc::new(InMemorySessionStorage::new()))
.sse(true) .stream_config(custom_stream_config.clone())
.build()
.await
.expect("Server should build successfully");
let handler = server
.handler()
.await
.expect("Handler should be created from server");
assert!(handler.sse_enabled, "SSE should be enabled");
let stream_manager = handler.get_stream_manager();
let actual_config = stream_manager.get_config();
assert_eq!(
actual_config.channel_buffer_size, custom_stream_config.channel_buffer_size,
"Custom channel_buffer_size should be preserved through builder → server → handler chain"
);
assert_eq!(
actual_config.max_replay_events, custom_stream_config.max_replay_events,
"Custom max_replay_events should be preserved through builder → server → handler chain"
);
assert_eq!(
actual_config.keepalive_interval_seconds,
custom_stream_config.keepalive_interval_seconds,
"Custom keepalive_interval_seconds should be preserved through builder → server → handler chain"
);
assert_eq!(
actual_config.cors_origin, custom_stream_config.cors_origin,
"Custom cors_origin should be preserved through builder → server → handler chain"
);
assert!(
Arc::strong_count(stream_manager) >= 1,
"Stream manager should be properly initialized"
);
let test_session_id = uuid::Uuid::now_v7().as_simple().to_string();
let subscriptions = stream_manager.get_subscriptions(&test_session_id).await;
assert!(
subscriptions.is_empty(),
"New session should have no subscriptions initially"
);
assert_eq!(
stream_manager.get_config().channel_buffer_size,
2048,
"Stream manager should be using the custom buffer size functionally"
);
}
#[tokio::test]
async fn test_non_streaming_runtime_sse_false() {
use crate::LambdaMcpServerBuilder;
use turul_mcp_session_storage::InMemorySessionStorage;
let server = LambdaMcpServerBuilder::new()
.name("test-non-streaming-sse-false")
.version("1.0.0")
.storage(Arc::new(InMemorySessionStorage::new()))
.sse(false) .build()
.await
.expect("Server should build successfully");
let handler = server
.handler()
.await
.expect("Handler should be created from server");
assert!(!handler.sse_enabled, "SSE should be disabled");
let lambda_req = Request::builder()
.method("POST")
.uri("/mcp")
.body(LambdaBody::Text(
r#"{"jsonrpc":"2.0","method":"initialize","id":1}"#.to_string(),
))
.unwrap();
let result = handler.handle(lambda_req).await;
assert!(
result.is_ok(),
"POST /mcp should work with non-streaming + sse(false)"
);
}
#[tokio::test]
async fn test_non_streaming_runtime_sse_true() {
use crate::LambdaMcpServerBuilder;
use turul_mcp_session_storage::InMemorySessionStorage;
let server = LambdaMcpServerBuilder::new()
.name("test-non-streaming-sse-true")
.version("1.0.0")
.storage(Arc::new(InMemorySessionStorage::new()))
.sse(true) .build()
.await
.expect("Server should build successfully");
let handler = server
.handler()
.await
.expect("Handler should be created from server");
assert!(handler.sse_enabled, "SSE should be enabled");
let lambda_req = Request::builder()
.method("POST")
.uri("/mcp")
.body(LambdaBody::Text(
r#"{"jsonrpc":"2.0","method":"initialize","id":1}"#.to_string(),
))
.unwrap();
let result = handler.handle(lambda_req).await;
assert!(
result.is_ok(),
"POST /mcp should work with non-streaming + sse(true)"
);
}
#[tokio::test]
async fn test_streaming_runtime_sse_false() {
use crate::LambdaMcpServerBuilder;
use turul_mcp_session_storage::InMemorySessionStorage;
let server = LambdaMcpServerBuilder::new()
.name("test-streaming-sse-false")
.version("1.0.0")
.storage(Arc::new(InMemorySessionStorage::new()))
.sse(false) .build()
.await
.expect("Server should build successfully");
let handler = server
.handler()
.await
.expect("Handler should be created from server");
assert!(!handler.sse_enabled, "SSE should be disabled");
let lambda_req = Request::builder()
.method("POST")
.uri("/mcp")
.body(LambdaBody::Text(
r#"{"jsonrpc":"2.0","method":"initialize","id":1}"#.to_string(),
))
.unwrap();
let result = handler.handle_streaming(lambda_req).await;
assert!(
result.is_ok(),
"Streaming runtime should work with sse(false)"
);
}
#[tokio::test]
async fn test_streaming_runtime_sse_true() {
use crate::LambdaMcpServerBuilder;
use turul_mcp_session_storage::InMemorySessionStorage;
let server = LambdaMcpServerBuilder::new()
.name("test-streaming-sse-true")
.version("1.0.0")
.storage(Arc::new(InMemorySessionStorage::new()))
.sse(true) .build()
.await
.expect("Server should build successfully");
let handler = server
.handler()
.await
.expect("Handler should be created from server");
assert!(handler.sse_enabled, "SSE should be enabled");
let lambda_req = Request::builder()
.method("POST")
.uri("/mcp")
.body(LambdaBody::Text(
r#"{"jsonrpc":"2.0","method":"initialize","id":1}"#.to_string(),
))
.unwrap();
let result = handler.handle_streaming(lambda_req).await;
assert!(
result.is_ok(),
"Streaming runtime should work with sse(true) for real-time streaming"
);
}
async fn build_strict_streaming_handler() -> LambdaMcpHandler {
use crate::LambdaMcpServerBuilder;
use turul_mcp_session_storage::InMemorySessionStorage;
let server = LambdaMcpServerBuilder::new()
.name("lifecycle-test")
.version("1.0.0")
.tool(LifecycleTestTool)
.storage(Arc::new(InMemorySessionStorage::new()))
.strict_lifecycle(true) .sse(true)
.build()
.await
.expect("build should succeed");
server.handler().await.expect("handler should succeed")
}
#[derive(Clone, Default)]
struct LifecycleTestTool;
impl turul_mcp_builders::traits::HasBaseMetadata for LifecycleTestTool {
fn name(&self) -> &str {
"ping_tool"
}
}
impl turul_mcp_builders::traits::HasDescription for LifecycleTestTool {
fn description(&self) -> Option<&str> {
Some("test tool")
}
}
impl turul_mcp_builders::traits::HasInputSchema for LifecycleTestTool {
fn input_schema(&self) -> &turul_mcp_protocol::ToolSchema {
static SCHEMA: std::sync::OnceLock<turul_mcp_protocol::ToolSchema> =
std::sync::OnceLock::new();
SCHEMA.get_or_init(turul_mcp_protocol::ToolSchema::object)
}
}
impl turul_mcp_builders::traits::HasOutputSchema for LifecycleTestTool {
fn output_schema(&self) -> Option<&turul_mcp_protocol::ToolSchema> {
None
}
}
impl turul_mcp_builders::traits::HasAnnotations for LifecycleTestTool {
fn annotations(&self) -> Option<&turul_mcp_protocol::tools::ToolAnnotations> {
None
}
}
impl turul_mcp_builders::traits::HasToolMeta for LifecycleTestTool {
fn tool_meta(&self) -> Option<&std::collections::HashMap<String, serde_json::Value>> {
None
}
}
impl turul_mcp_builders::traits::HasIcons for LifecycleTestTool {}
impl turul_mcp_builders::traits::HasExecution for LifecycleTestTool {}
#[async_trait::async_trait]
impl turul_mcp_server::McpTool for LifecycleTestTool {
async fn call(
&self,
_args: serde_json::Value,
_session: Option<turul_mcp_server::SessionContext>,
) -> turul_mcp_server::McpResult<turul_mcp_protocol::tools::CallToolResult> {
Ok(turul_mcp_protocol::tools::CallToolResult::success(vec![
turul_mcp_protocol::tools::ToolResult::text("pong"),
]))
}
}
fn streaming_mcp_request(body: &str, session_id: Option<&str>) -> LambdaRequest {
let mut builder = Request::builder()
.method("POST")
.uri("/mcp")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("MCP-Protocol-Version", "2025-11-25");
if let Some(sid) = session_id {
builder = builder.header("Mcp-Session-Id", sid);
}
builder.body(LambdaBody::Text(body.to_string())).unwrap()
}
async fn collect_streaming_body(
response: lambda_http::Response<
http_body_util::combinators::UnsyncBoxBody<bytes::Bytes, hyper::Error>,
>,
) -> (http::StatusCode, String) {
use http_body_util::BodyExt;
let status = response.status();
let session_id = response
.headers()
.get("Mcp-Session-Id")
.and_then(|v| v.to_str().ok())
.map(String::from);
let body_bytes = response
.into_body()
.collect()
.await
.map(|c| c.to_bytes())
.unwrap_or_default();
let body_str = String::from_utf8_lossy(&body_bytes).to_string();
let _ = session_id; (status, body_str)
}
fn extract_session_id(
response: &lambda_http::Response<
http_body_util::combinators::UnsyncBoxBody<bytes::Bytes, hyper::Error>,
>,
) -> Option<String> {
response
.headers()
.get("Mcp-Session-Id")
.and_then(|v| v.to_str().ok())
.map(String::from)
}
fn parse_response_json(body: &str) -> serde_json::Value {
let json_str = body
.lines()
.find(|line| line.starts_with("data: "))
.map(|line| &line[6..])
.unwrap_or(body.trim());
serde_json::from_str(json_str)
.unwrap_or_else(|e| panic!("Failed to parse JSON from body: {e}\nBody: {body}"))
}
#[tokio::test]
async fn test_lambda_streaming_strict_handshake_succeeds() {
let handler = build_strict_streaming_handler().await;
let init_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0", "method": "initialize", "id": 1,
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0.0" }
}
})
.to_string(),
None,
);
let init_resp = handler
.handle_streaming(init_req)
.await
.expect("initialize should succeed");
let session_id = extract_session_id(&init_resp).expect("must return session ID");
let (status, _body) = collect_streaming_body(init_resp).await;
assert_eq!(status, 200, "initialize should return 200");
let notif_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/initialized",
"params": {}
})
.to_string(),
Some(&session_id),
);
let notif_resp = handler
.handle_streaming(notif_req)
.await
.expect("notification should succeed");
let (status, _) = collect_streaming_body(notif_resp).await;
assert_eq!(status, 202, "notifications/initialized should return 202");
let list_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0", "method": "tools/list", "id": 2
})
.to_string(),
Some(&session_id),
);
let list_resp = handler
.handle_streaming(list_req)
.await
.expect("tools/list should succeed");
let (status, body) = collect_streaming_body(list_resp).await;
assert_eq!(status, 200, "tools/list should return 200");
let json = parse_response_json(&body);
assert!(
json["result"]["tools"].is_array(),
"tools/list should return tools array: {json}"
);
let call_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0", "method": "tools/call", "id": 3,
"params": { "name": "ping_tool", "arguments": {} }
})
.to_string(),
Some(&session_id),
);
let call_resp = handler
.handle_streaming(call_req)
.await
.expect("tools/call should succeed");
let (status, body) = collect_streaming_body(call_resp).await;
assert_eq!(status, 200, "tools/call should return 200");
let json = parse_response_json(&body);
assert!(
json["result"].is_object(),
"tools/call should return result: {json}"
);
}
#[tokio::test]
async fn test_lambda_streaming_strict_rejects_before_initialized() {
let handler = build_strict_streaming_handler().await;
let init_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0", "method": "initialize", "id": 1,
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0.0" }
}
})
.to_string(),
None,
);
let init_resp = handler.handle_streaming(init_req).await.unwrap();
let session_id = extract_session_id(&init_resp).unwrap();
let _ = collect_streaming_body(init_resp).await;
let list_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0", "method": "tools/list", "id": 2
})
.to_string(),
Some(&session_id),
);
let list_resp = handler.handle_streaming(list_req).await.unwrap();
let (_, body) = collect_streaming_body(list_resp).await;
let json = parse_response_json(&body);
assert!(
json["error"].is_object(),
"tools/list should return JSON-RPC error: {json}"
);
assert_eq!(
json["error"]["code"].as_i64().unwrap(),
-32031,
"tools/list must return SessionError code -32031, got: {json}"
);
assert!(
json["error"]["message"]
.as_str()
.unwrap()
.contains("notifications/initialized"),
"Error must mention notifications/initialized: {}",
json["error"]["message"]
);
let call_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0", "method": "tools/call", "id": 3,
"params": { "name": "ping_tool", "arguments": {} }
})
.to_string(),
Some(&session_id),
);
let call_resp = handler.handle_streaming(call_req).await.unwrap();
let (_, body) = collect_streaming_body(call_resp).await;
let json = parse_response_json(&body);
assert!(
json["error"].is_object(),
"tools/call should return JSON-RPC error: {json}"
);
assert_eq!(
json["error"]["code"].as_i64().unwrap(),
-32031,
"tools/call must return SessionError code -32031, got: {json}"
);
assert!(
json["error"]["message"]
.as_str()
.unwrap()
.contains("notifications/initialized"),
"Error must mention notifications/initialized: {}",
json["error"]["message"]
);
}
#[tokio::test]
async fn test_lambda_streaming_initialized_is_effective_immediately() {
let handler = build_strict_streaming_handler().await;
let init_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0", "method": "initialize", "id": 1,
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0.0" }
}
})
.to_string(),
None,
);
let init_resp = handler.handle_streaming(init_req).await.unwrap();
let session_id = extract_session_id(&init_resp).unwrap();
let _ = collect_streaming_body(init_resp).await;
let notif_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/initialized",
"params": {}
})
.to_string(),
Some(&session_id),
);
let notif_resp = handler.handle_streaming(notif_req).await.unwrap();
let (status, _) = collect_streaming_body(notif_resp).await;
assert_eq!(status, 202);
let list_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0", "method": "tools/list", "id": 2
})
.to_string(),
Some(&session_id),
);
let list_resp = handler.handle_streaming(list_req).await.unwrap();
let (status, body) = collect_streaming_body(list_resp).await;
assert_eq!(
status, 200,
"tools/list must succeed immediately after initialized"
);
let json = parse_response_json(&body);
assert!(
json["result"]["tools"].is_array(),
"Must return tools list, not error: {json}"
);
}
#[tokio::test]
async fn test_lambda_streaming_lenient_mode_allows_without_initialized() {
use crate::LambdaMcpServerBuilder;
use turul_mcp_session_storage::InMemorySessionStorage;
let server = LambdaMcpServerBuilder::new()
.name("lenient-test")
.version("1.0.0")
.tool(LifecycleTestTool)
.storage(Arc::new(InMemorySessionStorage::new()))
.strict_lifecycle(false) .sse(true)
.build()
.await
.unwrap();
let handler = server.handler().await.unwrap();
let init_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0", "method": "initialize", "id": 1,
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0.0" }
}
})
.to_string(),
None,
);
let init_resp = handler.handle_streaming(init_req).await.unwrap();
let session_id = extract_session_id(&init_resp).unwrap();
let _ = collect_streaming_body(init_resp).await;
let list_req = streaming_mcp_request(
&serde_json::json!({
"jsonrpc": "2.0", "method": "tools/list", "id": 2
})
.to_string(),
Some(&session_id),
);
let list_resp = handler.handle_streaming(list_req).await.unwrap();
let (status, body) = collect_streaming_body(list_resp).await;
assert_eq!(
status, 200,
"Lenient mode should allow tools/list without initialized"
);
let json = parse_response_json(&body);
assert!(
json["result"]["tools"].is_array(),
"Must return tools list in lenient mode: {json}"
);
}
#[cfg(feature = "cors")]
mod cors_streaming_routes {
use super::*;
use async_trait::async_trait;
use bytes::Bytes;
use http_body_util::Full;
use hyper::{Request as HyperRequest, Response as HyperResponse, StatusCode};
use turul_http_mcp_server::middleware::MiddlewareStack;
use turul_http_mcp_server::{
RouteBody, RouteHandler, RouteRegistry, StreamConfig, StreamManager,
};
struct StubRoute {
status: StatusCode,
body: &'static str,
}
#[async_trait]
impl RouteHandler for StubRoute {
async fn handle(&self, _req: HyperRequest<RouteBody>) -> HyperResponse<RouteBody> {
use http_body_util::BodyExt;
HyperResponse::builder()
.status(self.status)
.header("Content-Type", "application/json")
.body(
Full::new(Bytes::from(self.body))
.map_err(|never| match never {})
.boxed_unsync(),
)
.unwrap()
}
}
fn handler_with_route_and_cors(
registry: Arc<RouteRegistry>,
cors: Option<CorsConfig>,
) -> LambdaMcpHandler {
let session_storage = Arc::new(InMemorySessionStorage::new());
let stream_manager = Arc::new(StreamManager::new(session_storage.clone()));
let dispatcher = Arc::new(JsonRpcDispatcher::new());
let config = ServerConfig::default();
let capabilities = ServerCapabilities::default();
let middleware_stack = Arc::new(MiddlewareStack::new());
let handler = LambdaMcpHandler::with_middleware(
config,
dispatcher,
session_storage,
stream_manager,
StreamConfig::default(),
capabilities,
middleware_stack,
false,
registry,
);
match cors {
Some(cfg) => handler.with_cors(cfg),
None => handler,
}
}
fn get_request(path: &str, origin: &str) -> LambdaRequest {
Request::builder()
.method("GET")
.uri(path)
.header("Origin", origin)
.body(LambdaBody::Empty)
.unwrap()
}
#[tokio::test]
async fn streaming_custom_route_match_injects_cors() {
let mut registry = RouteRegistry::new();
registry.add_route(
"/.well-known/oauth-protected-resource",
Arc::new(StubRoute {
status: StatusCode::OK,
body: r#"{"resource":"https://example.test/mcp"}"#,
}),
);
let handler = handler_with_route_and_cors(
Arc::new(registry),
Some(CorsConfig::default()),
);
let req = get_request(
"/.well-known/oauth-protected-resource",
"https://client.example.test",
);
let resp = handler.handle_streaming(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert!(
resp.headers().contains_key("access-control-allow-origin"),
"matched streaming route must carry CORS headers",
);
assert!(
resp.headers().contains_key("access-control-expose-headers"),
"matched streaming route must expose configured headers",
);
}
#[tokio::test]
async fn streaming_route_validation_error_injects_cors() {
let registry = Arc::new({
let mut r = RouteRegistry::new();
r.add_route(
"/.well-known/oauth-protected-resource",
Arc::new(StubRoute {
status: StatusCode::OK,
body: "{}",
}),
);
r
});
let handler = handler_with_route_and_cors(registry, Some(CorsConfig::default()));
let req = get_request("/../etc/passwd", "https://client.example.test");
let resp = handler.handle_streaming(req).await.unwrap();
assert!(
resp.status().is_client_error(),
"path-traversal must be a 4xx, got {}",
resp.status(),
);
assert!(
resp.headers().contains_key("access-control-allow-origin"),
"validation-error streaming route must carry CORS headers",
);
}
#[tokio::test]
async fn streaming_custom_route_without_cors_config_returns_untouched() {
let mut registry = RouteRegistry::new();
registry.add_route(
"/.well-known/oauth-protected-resource",
Arc::new(StubRoute {
status: StatusCode::OK,
body: "{}",
}),
);
let handler = handler_with_route_and_cors(Arc::new(registry), None);
let req = get_request(
"/.well-known/oauth-protected-resource",
"https://client.example.test",
);
let resp = handler.handle_streaming(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert!(
!resp.headers().contains_key("access-control-allow-origin"),
"no CORS config → no CORS headers (got {:?})",
resp.headers(),
);
}
}
#[cfg(feature = "cors")]
mod cors_streaming_oauth {
use super::*;
use async_trait::async_trait;
use turul_http_mcp_server::middleware::{
DispatcherResult, McpMiddleware, MiddlewareError, MiddlewareStack, RequestContext,
SessionInjection,
};
use turul_http_mcp_server::{StreamConfig, StreamManager};
use turul_mcp_session_storage::SessionView;
struct ForceChallenge;
#[async_trait]
impl McpMiddleware for ForceChallenge {
fn runs_before_session(&self) -> bool {
true
}
async fn before_dispatch(
&self,
_ctx: &mut RequestContext<'_>,
_session: Option<&dyn SessionView>,
_injection: &mut SessionInjection,
) -> std::result::Result<(), MiddlewareError> {
Err(MiddlewareError::http_challenge(
401,
"Bearer realm=\"mcp\", resource_metadata=\"https://example.test/.well-known/oauth-protected-resource\"",
))
}
async fn after_dispatch(
&self,
_ctx: &RequestContext<'_>,
_result: &mut DispatcherResult,
) -> std::result::Result<(), MiddlewareError> {
Ok(())
}
}
#[tokio::test]
async fn streaming_401_challenge_has_cors_and_exposes_www_authenticate() {
let session_storage = Arc::new(InMemorySessionStorage::new());
let stream_manager = Arc::new(StreamManager::new(session_storage.clone()));
let dispatcher = Arc::new(JsonRpcDispatcher::new());
let config = ServerConfig::default();
let capabilities = ServerCapabilities::default();
let mut middleware = MiddlewareStack::new();
middleware.push(Arc::new(ForceChallenge));
let middleware = Arc::new(middleware);
let route_registry =
Arc::new(turul_http_mcp_server::RouteRegistry::new());
let handler = LambdaMcpHandler::with_middleware(
config,
dispatcher,
session_storage,
stream_manager,
StreamConfig::default(),
capabilities,
middleware,
false,
route_registry,
)
.with_cors(CorsConfig::default());
let req = Request::builder()
.method("POST")
.uri("/mcp")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("MCP-Protocol-Version", "2025-11-25")
.header("Origin", "https://client.example.test")
.body(LambdaBody::Text(
r#"{"jsonrpc":"2.0","method":"initialize","id":1}"#.to_string(),
))
.unwrap();
let resp = handler.handle_streaming(req).await.unwrap();
let headers = resp.headers();
assert_eq!(resp.status(), 401, "challenge must be 401");
assert!(
headers.contains_key("www-authenticate"),
"WWW-Authenticate must be preserved through streaming transport",
);
assert!(
headers.contains_key("access-control-allow-origin"),
"401 response must carry Access-Control-Allow-Origin",
);
let expose = headers
.get("access-control-expose-headers")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
assert!(
expose
.split(',')
.map(str::trim)
.any(|h| h.eq_ignore_ascii_case("WWW-Authenticate")),
"expose-headers must include WWW-Authenticate; got {expose:?}",
);
}
}
}