mod common;
use actix_web::{App, HttpServer};
use common::headers_test_service::HeadersTestService;
use futures::StreamExt;
use reqwest::Response;
use rmcp::transport::streamable_http_server::session::local::LocalSessionManager;
use rmcp_actix_web::transport::StreamableHttpService;
use serde_json::{Value, json};
use std::sync::Arc;
use std::time::Duration;
async fn extract_auth_from_sse_response(response: Response) -> Option<String> {
let mut body = Vec::new();
let mut stream = response.bytes_stream();
tokio::time::timeout(Duration::from_secs(2), async {
loop {
let body_str = String::from_utf8_lossy(&body);
for line in body_str.lines() {
if let Some(json_str) = line.strip_prefix("data: ")
&& let Ok(response_json) = serde_json::from_str::<Value>(json_str)
&& let Some(text_value) = response_json.pointer("/result/content/0/text")
&& let Some(text_str) = text_value.as_str()
&& let Ok(auth_response) = serde_json::from_str::<Value>(text_str)
&& let Some(auth) = auth_response.get("authorization")
{
return auth.as_str().map(String::from);
}
}
if body.len() > 4096 {
return None;
}
match stream.next().await {
Some(Ok(bytes)) => body.extend_from_slice(&bytes),
_ => return None,
}
}
})
.await
.ok()
.flatten()
}
#[cfg(feature = "authorization-token-passthrough")]
#[actix_web::test]
async fn test_authorization_forwarded_in_streamable_http_stateless() {
let _ = tracing_subscriber::fmt()
.with_env_filter("rmcp_actix_web=debug")
.with_test_writer()
.try_init();
let service = StreamableHttpService::builder()
.service_factory(Arc::new(|| Ok(HeadersTestService::new())))
.session_manager(Arc::new(LocalSessionManager::default()))
.stateful_mode(false)
.build();
let server = HttpServer::new(move || {
App::new().service(actix_web::web::scope("/mcp").service(service.clone().scope()))
})
.bind("127.0.0.1:0")
.expect("Failed to bind server");
let addr = *server.addrs().first().unwrap();
let server_handle = server.run();
let server_task = tokio::spawn(async move {
let _ = server_handle.await;
});
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let client = reqwest::Client::new();
let url = format!("http://{}/mcp", addr);
let init_request = json!({
"jsonrpc": "2.0",
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
},
"id": 1
});
let response = client
.post(&url)
.header("Authorization", "Bearer test-token-abc123")
.header("X-Custom-Header", "should-not-be-forwarded")
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&init_request)
.send()
.await
.expect("Failed to send request");
assert_eq!(response.status(), 200);
let mut body = Vec::new();
let mut stream = response.bytes_stream();
let _ = tokio::time::timeout(Duration::from_secs(2), async {
while let Some(chunk) = stream.next().await {
if let Ok(bytes) = chunk {
body.extend_from_slice(&bytes);
if body.ends_with(b"\n\n") || body.len() > 4096 {
break;
}
}
}
})
.await;
let body_str = String::from_utf8_lossy(&body);
assert!(
body_str.contains("data: "),
"Response should be in SSE format"
);
let tool_request = json!({
"jsonrpc": "2.0",
"method": "tools/call",
"params": {
"name": "get_headers"
},
"id": 2
});
let tool_response = client
.post(&url)
.header("Authorization", "Bearer test-token-abc123")
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&tool_request)
.send()
.await
.expect("Failed to send tool request");
assert_eq!(tool_response.status(), 200);
server_task.abort();
}
#[cfg(feature = "authorization-token-passthrough")]
#[actix_web::test]
async fn test_authorization_forwarded_in_streamable_http_stateful() {
let _ = tracing_subscriber::fmt()
.with_env_filter("rmcp_actix_web=debug")
.with_test_writer()
.try_init();
let service = StreamableHttpService::builder()
.service_factory(Arc::new(|| Ok(HeadersTestService::new())))
.session_manager(Arc::new(LocalSessionManager::default()))
.stateful_mode(true)
.build();
let server = HttpServer::new(move || {
App::new().service(actix_web::web::scope("/mcp").service(service.clone().scope()))
})
.bind("127.0.0.1:0")
.expect("Failed to bind server");
let addr = *server.addrs().first().unwrap();
let server_handle = server.run();
let server_task = tokio::spawn(async move {
let _ = server_handle.await;
});
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let client = reqwest::Client::new();
let url = format!("http://{}/mcp", addr);
let init_request = json!({
"jsonrpc": "2.0",
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
},
"id": 1
});
let response = client
.post(&url)
.header("Authorization", "Bearer initial-token")
.header("X-Another-Header", "should-not-be-forwarded")
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&init_request)
.send()
.await
.expect("Failed to send request");
assert_eq!(response.status(), 200);
let session_id = response
.headers()
.get("Mcp-Session-Id")
.and_then(|v| v.to_str().ok())
.expect("Should have session ID")
.to_string();
assert!(!session_id.is_empty());
let mut init_body = Vec::new();
let mut init_stream = response.bytes_stream();
let _ = tokio::time::timeout(Duration::from_secs(2), async {
while let Some(chunk) = init_stream.next().await {
if let Ok(bytes) = chunk {
init_body.extend_from_slice(&bytes);
if init_body.ends_with(b"\n\n") || init_body.len() > 4096 {
break;
}
}
}
})
.await;
let initialized_notification = json!({
"jsonrpc": "2.0",
"method": "notifications/initialized"
});
let initialized_response = client
.post(&url)
.header("Authorization", "Bearer initial-token")
.header("Mcp-Session-Id", &session_id)
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&initialized_notification)
.send()
.await
.expect("Failed to send initialized notification");
assert_eq!(initialized_response.status(), 202);
let tool_request = json!({
"jsonrpc": "2.0",
"method": "tools/call",
"params": {
"name": "get_current_auth"
},
"id": 2
});
let tool_response = client
.post(&url)
.header("Authorization", "Bearer subsequent-token")
.header("Mcp-Session-Id", &session_id)
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&tool_request)
.send()
.await
.expect("Failed to send tool request");
assert_eq!(tool_response.status(), 200);
let auth = extract_auth_from_sse_response(tool_response).await;
assert_eq!(
auth,
Some("Bearer subsequent-token".to_string()),
"Bug #26: Authorization header should be forwarded for existing sessions"
);
let rotation_request = json!({
"jsonrpc": "2.0",
"method": "tools/call",
"params": {
"name": "get_current_auth"
},
"id": 3
});
let rotation_response = client
.post(&url)
.header("Authorization", "Bearer rotated-token")
.header("Mcp-Session-Id", &session_id)
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&rotation_request)
.send()
.await
.expect("Failed to send rotation request");
assert_eq!(rotation_response.status(), 200);
let rotated_auth = extract_auth_from_sse_response(rotation_response).await;
assert_eq!(
rotated_auth,
Some("Bearer rotated-token".to_string()),
"Token rotation should work within same session (OAuth 2.1 best practice)"
);
server_task.abort();
}
#[cfg(feature = "authorization-token-passthrough")]
#[actix_web::test]
async fn test_malformed_bearer_tokens_not_forwarded() {
let _ = tracing_subscriber::fmt()
.with_env_filter("rmcp_actix_web=debug")
.with_test_writer()
.try_init();
let service = StreamableHttpService::builder()
.service_factory(Arc::new(|| Ok(HeadersTestService::new())))
.session_manager(Arc::new(LocalSessionManager::default()))
.stateful_mode(true)
.build();
let server = HttpServer::new(move || {
App::new().service(actix_web::web::scope("/").service(service.clone().scope()))
})
.bind("127.0.0.1:0")
.expect("Failed to bind to port");
let port = server.addrs()[0].port();
let server_task = tokio::spawn(server.run());
let url = format!("http://127.0.0.1:{}", port);
tokio::time::sleep(Duration::from_millis(100)).await;
let client = reqwest::Client::new();
let init_request = json!({
"jsonrpc": "2.0",
"method": "initialize",
"params": {
"protocolVersion": "0.1.0",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
},
"id": 1
});
let response = client
.post(&url)
.header("Authorization", "Bearer") .header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&init_request)
.send()
.await
.expect("Failed to send request");
assert_eq!(response.status(), 200);
let session_id = response
.headers()
.get("Mcp-Session-Id")
.and_then(|v| v.to_str().ok())
.expect("Should have session ID")
.to_string();
let tool_request = json!({
"jsonrpc": "2.0",
"method": "tools/call",
"params": {
"name": "get_current_auth"
},
"id": 2
});
let tool_response = client
.post(&url)
.header("Authorization", "Bearer ") .header("Mcp-Session-Id", &session_id)
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&tool_request)
.send()
.await
.expect("Failed to send tool request");
assert_eq!(tool_response.status(), 200);
let auth = extract_auth_from_sse_response(tool_response).await;
assert_eq!(auth, None, "Malformed Bearer token should not be forwarded");
server_task.abort();
}
#[actix_web::test]
async fn test_non_bearer_authorization_not_forwarded() {
let _ = tracing_subscriber::fmt()
.with_env_filter("rmcp_actix_web=debug")
.with_test_writer()
.try_init();
let service = StreamableHttpService::builder()
.service_factory(Arc::new(|| Ok(HeadersTestService::new())))
.session_manager(Arc::new(LocalSessionManager::default()))
.stateful_mode(false)
.build();
let server = HttpServer::new(move || {
App::new().service(actix_web::web::scope("/mcp").service(service.clone().scope()))
})
.bind("127.0.0.1:0")
.expect("Failed to bind server");
let addr = *server.addrs().first().unwrap();
let server_handle = server.run();
let server_task = tokio::spawn(async move {
let _ = server_handle.await;
});
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let client = reqwest::Client::new();
let url = format!("http://{}/mcp", addr);
let init_request = json!({
"jsonrpc": "2.0",
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
},
"id": 1
});
let response = client
.post(&url)
.header("Authorization", "Basic dXNlcjpwYXNz")
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&init_request)
.send()
.await
.expect("Failed to send request");
assert_eq!(response.status(), 200);
let custom_auth_request = json!({
"jsonrpc": "2.0",
"method": "tools/call",
"params": {
"name": "get_current_auth"
},
"id": 2
});
let response = client
.post(&url)
.header("Authorization", "CustomScheme sometoken123")
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&custom_auth_request)
.send()
.await
.expect("Failed to send request");
assert_eq!(response.status(), 200);
let auth = extract_auth_from_sse_response(response).await;
assert_eq!(
auth, None,
"Non-Bearer authorization should not be forwarded"
);
server_task.abort();
}
#[actix_web::test]
async fn test_missing_authorization_doesnt_break_service() {
let _ = tracing_subscriber::fmt()
.with_env_filter("rmcp_actix_web=debug")
.with_test_writer()
.try_init();
let service = StreamableHttpService::builder()
.service_factory(Arc::new(|| Ok(HeadersTestService::new())))
.session_manager(Arc::new(LocalSessionManager::default()))
.stateful_mode(false)
.build();
let server = HttpServer::new(move || {
App::new().service(actix_web::web::scope("/mcp").service(service.clone().scope()))
})
.bind("127.0.0.1:0")
.expect("Failed to bind server");
let addr = *server.addrs().first().unwrap();
let server_handle = server.run();
let server_task = tokio::spawn(async move {
let _ = server_handle.await;
});
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let client = reqwest::Client::new();
let url = format!("http://{}/mcp", addr);
let init_request = json!({
"jsonrpc": "2.0",
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
},
"id": 1
});
let response = client
.post(&url)
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&init_request)
.send()
.await
.expect("Failed to send request");
assert_eq!(response.status(), 200);
let tool_request = json!({
"jsonrpc": "2.0",
"method": "tools/call",
"params": {
"name": "get_current_auth"
},
"id": 2
});
let tool_response = client
.post(&url)
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.json(&tool_request)
.send()
.await
.expect("Failed to send tool request");
assert_eq!(tool_response.status(), 200);
let auth = extract_auth_from_sse_response(tool_response).await;
assert_eq!(
auth, None,
"No authorization should be present when header is missing"
);
server_task.abort();
}