use std::collections::HashMap;
use std::str::FromStr;
use bytes::Bytes;
use http_body_util::{BodyExt, Full};
use hyper::Response as HyperResponse;
use lambda_http::{Body as LambdaBody, Request as LambdaRequest, Response as LambdaResponse};
use tracing::{debug, trace};
use crate::error::{LambdaError, Result};
type UnifiedMcpBody = http_body_util::combinators::UnsyncBoxBody<Bytes, hyper::Error>;
fn infallible_to_hyper_error(never: std::convert::Infallible) -> hyper::Error {
match never {}
}
type MappedFullBody =
http_body_util::combinators::MapErr<Full<Bytes>, fn(std::convert::Infallible) -> hyper::Error>;
pub fn lambda_to_hyper_request(
lambda_req: LambdaRequest,
) -> Result<hyper::Request<MappedFullBody>> {
let authorizer_fields = extract_authorizer_context(&lambda_req);
let (mut parts, lambda_body) = lambda_req.into_parts();
for (field_name, field_value) in authorizer_fields {
let header_name = format!("x-authorizer-{}", field_name);
let Ok(name) = http::HeaderName::from_str(&header_name) else {
debug!(
"Skipping authorizer field '{}' - invalid header name",
field_name
);
continue;
};
let Ok(value) = http::HeaderValue::from_str(&field_value) else {
debug!(
"Skipping authorizer field '{}' - invalid header value",
field_name
);
continue;
};
parts.headers.insert(name, value);
trace!(
"Injected authorizer header: {} = {}",
header_name, field_value
);
}
let body_bytes = match lambda_body {
LambdaBody::Empty => Bytes::new(),
LambdaBody::Text(s) => Bytes::from(s),
LambdaBody::Binary(b) => Bytes::from(b),
_ => Bytes::new(),
};
let full_body = Full::new(body_bytes)
.map_err(infallible_to_hyper_error as fn(std::convert::Infallible) -> hyper::Error);
let hyper_req = hyper::Request::from_parts(parts, full_body);
debug!(
"Converted Lambda request: {} {} -> hyper::Request<Full<Bytes>>",
hyper_req.method(),
hyper_req.uri()
);
Ok(hyper_req)
}
pub async fn hyper_to_lambda_response(
hyper_resp: HyperResponse<UnifiedMcpBody>,
) -> Result<LambdaResponse<LambdaBody>> {
let (parts, body) = hyper_resp.into_parts();
let body_bytes = match body.collect().await {
Ok(collected) => collected.to_bytes(),
Err(err) => {
return Err(LambdaError::Body(format!(
"Failed to collect response body: {}",
err
)));
}
};
let lambda_body = if body_bytes.is_empty() {
LambdaBody::Empty
} else {
match String::from_utf8(body_bytes.to_vec()) {
Ok(text) => LambdaBody::Text(text),
Err(_) => LambdaBody::Binary(body_bytes.to_vec()),
}
};
let lambda_resp = LambdaResponse::from_parts(parts, lambda_body);
debug!(
"Converted hyper response -> Lambda response (status: {})",
lambda_resp.status()
);
Ok(lambda_resp)
}
pub fn hyper_to_lambda_streaming(
hyper_resp: HyperResponse<UnifiedMcpBody>,
) -> lambda_http::Response<UnifiedMcpBody> {
let (parts, body) = hyper_resp.into_parts();
let lambda_resp = lambda_http::Response::from_parts(parts, body);
debug!(
"Converted hyper response -> Lambda streaming response (status: {})",
lambda_resp.status()
);
lambda_resp
}
pub fn camel_to_snake(s: &str) -> String {
let mut result = String::new();
let chars: Vec<char> = s.chars().collect();
for i in 0..chars.len() {
let ch = chars[i];
if ch.is_uppercase() {
let is_first = i == 0;
let prev_is_lower = i > 0 && chars[i - 1].is_lowercase();
let next_is_lower = i + 1 < chars.len() && chars[i + 1].is_lowercase();
if !is_first && (prev_is_lower || next_is_lower) {
result.push('_');
}
result.push(ch.to_ascii_lowercase());
} else {
result.push(ch);
}
}
result
}
pub fn sanitize_authorizer_field_name(field: &str) -> String {
let snake_case = camel_to_snake(field);
snake_case
.to_ascii_lowercase()
.chars()
.map(|c| {
if c.is_ascii_alphanumeric() || c == '_' || c == '-' {
c
} else {
'-'
}
})
.collect()
}
pub fn extract_authorizer_context(req: &LambdaRequest) -> HashMap<String, String> {
use lambda_http::request::RequestContext;
let mut fields = HashMap::new();
let Some(request_context) = req.extensions().get::<RequestContext>() else {
return fields; };
match request_context {
RequestContext::ApiGatewayV1(ctx) => {
debug!(
authorizer_field_count = ctx.authorizer.fields.len(),
authorizer_keys = ?ctx.authorizer.fields.keys().collect::<Vec<_>>(),
"V1 REST API authorizer context"
);
}
RequestContext::ApiGatewayV2(ctx) => {
if let Some(ref authorizer) = ctx.authorizer {
debug!(
authorizer_field_count = authorizer.fields.len(),
authorizer_keys = ?authorizer.fields.keys().collect::<Vec<_>>(),
"V2 HTTP API authorizer context"
);
} else {
debug!("V2 HTTP API: no authorizer present");
}
}
_ => {
debug!("Non-API Gateway request context (ALB or other)");
}
}
let mut authorizer_fields_map = HashMap::new();
match request_context {
RequestContext::ApiGatewayV2(ctx) => {
if let Some(ref authorizer) = ctx.authorizer {
for (key, value) in &authorizer.fields {
authorizer_fields_map.insert(key.clone(), value.clone());
}
}
}
RequestContext::ApiGatewayV1(ctx) => {
if let Some(serde_json::Value::Object(auth_map)) = ctx.authorizer.fields.get("lambda") {
for (key, value) in auth_map {
authorizer_fields_map.insert(key.clone(), value.clone());
}
} else {
for (key, value) in &ctx.authorizer.fields {
if key == "principalId"
|| key == "integrationLatency"
|| key == "usageIdentifierKey"
{
continue;
}
authorizer_fields_map.insert(key.clone(), value.clone());
}
}
}
_ => {} }
for (key, value) in authorizer_fields_map {
let sanitized_key = sanitize_authorizer_field_name(&key);
let value_str = match value {
serde_json::Value::String(s) => s,
other => other.to_string(), };
fields.insert(sanitized_key, value_str);
}
if !fields.is_empty() {
debug!(
"Extracted {} authorizer fields from Lambda context",
fields.len()
);
}
fields
}
pub fn extract_mcp_headers(req: &LambdaRequest) -> HashMap<String, String> {
let mut mcp_headers = HashMap::new();
if let Some(session_id) = req.headers().get("mcp-session-id")
&& let Ok(session_id_str) = session_id.to_str()
{
mcp_headers.insert("mcp-session-id".to_string(), session_id_str.to_string());
}
if let Some(protocol_version) = req.headers().get("mcp-protocol-version")
&& let Ok(version_str) = protocol_version.to_str()
{
mcp_headers.insert("mcp-protocol-version".to_string(), version_str.to_string());
}
if let Some(last_event_id) = req.headers().get("last-event-id")
&& let Ok(event_id_str) = last_event_id.to_str()
{
mcp_headers.insert("last-event-id".to_string(), event_id_str.to_string());
}
trace!("Extracted MCP headers: {:?}", mcp_headers);
mcp_headers
}
pub fn inject_mcp_headers(resp: &mut LambdaResponse<LambdaBody>, headers: HashMap<String, String>) {
for (name, value) in headers {
if let (Ok(header_name), Ok(header_value)) = (
http::HeaderName::from_bytes(name.as_bytes()),
http::HeaderValue::from_str(&value),
) {
resp.headers_mut().insert(header_name, header_value);
debug!("Injected MCP header: {} = {}", name, value);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use http::{HeaderValue, Method, Request, StatusCode};
use http_body_util::Full;
#[test]
fn test_lambda_to_hyper_request_conversion() {
let mut lambda_req = Request::builder()
.method(Method::POST)
.uri("/mcp")
.body(LambdaBody::Text(
r#"{"jsonrpc":"2.0","method":"initialize","id":1}"#.to_string(),
))
.unwrap();
let headers = lambda_req.headers_mut();
headers.insert("content-type", HeaderValue::from_static("application/json"));
headers.insert(
"mcp-session-id",
HeaderValue::from_static("test-session-123"),
);
headers.insert(
"mcp-protocol-version",
HeaderValue::from_static("2025-11-25"),
);
let hyper_req = lambda_to_hyper_request(lambda_req).unwrap();
assert_eq!(hyper_req.method(), &Method::POST);
assert_eq!(hyper_req.uri().path(), "/mcp");
assert_eq!(
hyper_req.headers().get("content-type").unwrap(),
"application/json"
);
assert_eq!(
hyper_req.headers().get("mcp-session-id").unwrap(),
"test-session-123"
);
assert_eq!(
hyper_req.headers().get("mcp-protocol-version").unwrap(),
"2025-11-25"
);
}
#[test]
fn test_lambda_to_hyper_empty_body() {
let lambda_req = Request::builder()
.method(Method::GET)
.uri("/sse")
.body(LambdaBody::Empty)
.unwrap();
let hyper_req = lambda_to_hyper_request(lambda_req).unwrap();
assert_eq!(hyper_req.method(), &Method::GET);
assert_eq!(hyper_req.uri().path(), "/sse");
}
#[test]
fn test_lambda_to_hyper_binary_body() {
let test_data = vec![0x48, 0x65, 0x6c, 0x6c, 0x6f]; let lambda_req = Request::builder()
.method(Method::POST)
.uri("/binary")
.body(LambdaBody::Binary(test_data.clone()))
.unwrap();
let hyper_req = lambda_to_hyper_request(lambda_req).unwrap();
assert_eq!(hyper_req.method(), &Method::POST);
assert_eq!(hyper_req.uri().path(), "/binary");
}
#[tokio::test]
async fn test_hyper_to_lambda_response_conversion() {
let json_body = r#"{"jsonrpc":"2.0","id":1,"result":{"capabilities":{}}}"#;
let full_body = Full::new(Bytes::from(json_body));
let boxed_body = full_body.map_err(|never| match never {}).boxed_unsync();
let hyper_resp = hyper::Response::builder()
.status(StatusCode::OK)
.header("content-type", "application/json")
.header("mcp-session-id", "resp-session-456")
.body(boxed_body)
.unwrap();
let lambda_resp = hyper_to_lambda_response(hyper_resp).await.unwrap();
assert_eq!(lambda_resp.status(), StatusCode::OK);
assert_eq!(
lambda_resp.headers().get("content-type").unwrap(),
"application/json"
);
assert_eq!(
lambda_resp.headers().get("mcp-session-id").unwrap(),
"resp-session-456"
);
match lambda_resp.body() {
LambdaBody::Text(text) => assert_eq!(text, json_body),
_ => panic!("Expected text body"),
}
}
#[tokio::test]
async fn test_hyper_to_lambda_empty_response() {
let empty_body = Full::new(Bytes::new());
let boxed_body = empty_body.map_err(|never| match never {}).boxed_unsync();
let hyper_resp = hyper::Response::builder()
.status(StatusCode::NO_CONTENT)
.body(boxed_body)
.unwrap();
let lambda_resp = hyper_to_lambda_response(hyper_resp).await.unwrap();
assert_eq!(lambda_resp.status(), StatusCode::NO_CONTENT);
match lambda_resp.body() {
LambdaBody::Empty => {} _ => panic!("Expected empty body"),
}
}
#[test]
fn test_hyper_to_lambda_streaming() {
let stream_body = Full::new(Bytes::from("data: test\n\n"));
let boxed_body = stream_body.map_err(|never| match never {}).boxed_unsync();
let hyper_resp = hyper::Response::builder()
.status(StatusCode::OK)
.header("content-type", "text/event-stream")
.header("cache-control", "no-cache")
.body(boxed_body)
.unwrap();
let lambda_resp = hyper_to_lambda_streaming(hyper_resp);
assert_eq!(lambda_resp.status(), StatusCode::OK);
assert_eq!(
lambda_resp.headers().get("content-type").unwrap(),
"text/event-stream"
);
assert_eq!(
lambda_resp.headers().get("cache-control").unwrap(),
"no-cache"
);
}
#[tokio::test]
async fn test_mcp_headers_extraction() {
use http::{HeaderValue, Request};
let mut request = Request::builder()
.method("POST")
.uri("/mcp")
.body(LambdaBody::Empty)
.unwrap();
let headers = request.headers_mut();
headers.insert("mcp-session-id", HeaderValue::from_static("sess-123"));
headers.insert(
"mcp-protocol-version",
HeaderValue::from_static("2025-11-25"),
);
headers.insert("last-event-id", HeaderValue::from_static("event-456"));
let mcp_headers = extract_mcp_headers(&request);
assert_eq!(
mcp_headers.get("mcp-session-id"),
Some(&"sess-123".to_string())
);
assert_eq!(
mcp_headers.get("mcp-protocol-version"),
Some(&"2025-11-25".to_string())
);
assert_eq!(
mcp_headers.get("last-event-id"),
Some(&"event-456".to_string())
);
}
#[tokio::test]
async fn test_mcp_headers_injection() {
use lambda_http::Body;
let mut lambda_resp = LambdaResponse::builder()
.status(200)
.body(Body::Empty)
.unwrap();
let mut headers = HashMap::new();
headers.insert("mcp-session-id".to_string(), "sess-789".to_string());
headers.insert("mcp-protocol-version".to_string(), "2025-11-25".to_string());
inject_mcp_headers(&mut lambda_resp, headers);
assert_eq!(
lambda_resp.headers().get("mcp-session-id").unwrap(),
"sess-789"
);
assert_eq!(
lambda_resp.headers().get("mcp-protocol-version").unwrap(),
"2025-11-25"
);
}
mod authorizer_tests {
use super::*;
#[test]
fn test_sanitize_field_name_camelcase() {
assert_eq!(sanitize_authorizer_field_name("accountId"), "account_id");
assert_eq!(sanitize_authorizer_field_name("entityType"), "entity_type");
assert_eq!(sanitize_authorizer_field_name("deviceId"), "device_id");
assert_eq!(sanitize_authorizer_field_name("userId"), "user_id");
assert_eq!(sanitize_authorizer_field_name("tenantId"), "tenant_id");
assert_eq!(
sanitize_authorizer_field_name("customClaim"),
"custom_claim"
);
}
#[test]
fn test_sanitize_field_name_snake_case() {
assert_eq!(sanitize_authorizer_field_name("device_id"), "device_id");
assert_eq!(sanitize_authorizer_field_name("user_name"), "user_name");
assert_eq!(sanitize_authorizer_field_name("tenant_id"), "tenant_id");
}
#[test]
fn test_sanitize_field_name_acronyms() {
assert_eq!(sanitize_authorizer_field_name("APIKey"), "api_key");
assert_eq!(
sanitize_authorizer_field_name("HTTPSEnabled"),
"https_enabled"
);
assert_eq!(sanitize_authorizer_field_name("XMLParser"), "xml_parser");
}
#[test]
fn test_sanitize_field_name_with_numbers() {
assert_eq!(sanitize_authorizer_field_name("userId123"), "user_id123");
assert_eq!(sanitize_authorizer_field_name("device2Id"), "device2_id");
}
#[test]
fn test_sanitize_field_name_special_chars() {
assert_eq!(sanitize_authorizer_field_name("user@email"), "user-email");
assert_eq!(sanitize_authorizer_field_name("test.field"), "test-field");
assert_eq!(sanitize_authorizer_field_name("a/b/c"), "a-b-c");
}
#[test]
fn test_sanitize_field_name_unicode() {
assert_eq!(sanitize_authorizer_field_name("用户"), "--");
}
#[test]
fn test_extract_authorizer_no_context() {
let lambda_req = Request::builder()
.method(Method::POST)
.uri("/mcp")
.body(LambdaBody::Empty)
.unwrap();
let fields = extract_authorizer_context(&lambda_req);
assert!(fields.is_empty());
}
#[test]
fn test_lambda_to_hyper_without_authorizer() {
let lambda_req = Request::builder()
.method(Method::POST)
.uri("/mcp")
.header("content-type", "application/json")
.body(LambdaBody::Empty)
.unwrap();
let hyper_req = lambda_to_hyper_request(lambda_req).unwrap();
assert!(hyper_req.headers().get("x-authorizer-account_id").is_none());
assert_eq!(
hyper_req.headers().get("content-type").unwrap(),
"application/json"
);
}
fn request_with_context(ctx: lambda_http::request::RequestContext) -> LambdaRequest {
let mut req = Request::builder()
.method(Method::POST)
.uri("/mcp")
.body(LambdaBody::Empty)
.unwrap();
req.extensions_mut().insert(ctx);
req
}
#[test]
fn test_extract_authorizer_v1_top_level_fields() {
use aws_lambda_events::apigw::{
ApiGatewayProxyRequestContext, ApiGatewayRequestAuthorizer,
};
let mut authorizer = ApiGatewayRequestAuthorizer::default();
authorizer
.fields
.insert("userId".to_string(), serde_json::json!("user-123"));
authorizer
.fields
.insert("tenantId".to_string(), serde_json::json!("tenant-456"));
authorizer
.fields
.insert("role".to_string(), serde_json::json!("admin"));
let mut v1_ctx = ApiGatewayProxyRequestContext::default();
v1_ctx.authorizer = authorizer;
let req =
request_with_context(lambda_http::request::RequestContext::ApiGatewayV1(v1_ctx));
let fields = extract_authorizer_context(&req);
assert_eq!(fields.get("user_id"), Some(&"user-123".to_string()));
assert_eq!(fields.get("tenant_id"), Some(&"tenant-456".to_string()));
assert_eq!(fields.get("role"), Some(&"admin".to_string()));
}
#[test]
fn test_extract_authorizer_v1_nested_lambda() {
use aws_lambda_events::apigw::{
ApiGatewayProxyRequestContext, ApiGatewayRequestAuthorizer,
};
let mut authorizer = ApiGatewayRequestAuthorizer::default();
authorizer.fields.insert(
"lambda".to_string(),
serde_json::json!({
"userId": "user-123",
"tenantId": "tenant-456"
}),
);
let mut v1_ctx = ApiGatewayProxyRequestContext::default();
v1_ctx.authorizer = authorizer;
let req =
request_with_context(lambda_http::request::RequestContext::ApiGatewayV1(v1_ctx));
let fields = extract_authorizer_context(&req);
assert_eq!(fields.get("user_id"), Some(&"user-123".to_string()));
assert_eq!(fields.get("tenant_id"), Some(&"tenant-456".to_string()));
}
#[test]
fn test_extract_authorizer_v1_skips_internal_fields() {
use aws_lambda_events::apigw::{
ApiGatewayProxyRequestContext, ApiGatewayRequestAuthorizer,
};
let mut authorizer = ApiGatewayRequestAuthorizer::default();
authorizer
.fields
.insert("userId".to_string(), serde_json::json!("user-123"));
authorizer.fields.insert(
"principalId".to_string(),
serde_json::json!("principal-abc"),
);
authorizer
.fields
.insert("integrationLatency".to_string(), serde_json::json!(42));
authorizer.fields.insert(
"usageIdentifierKey".to_string(),
serde_json::json!("api-key-xyz"),
);
let mut v1_ctx = ApiGatewayProxyRequestContext::default();
v1_ctx.authorizer = authorizer;
let req =
request_with_context(lambda_http::request::RequestContext::ApiGatewayV1(v1_ctx));
let fields = extract_authorizer_context(&req);
assert_eq!(fields.get("user_id"), Some(&"user-123".to_string()));
assert!(
!fields.contains_key("principal_id"),
"principalId should be skipped"
);
assert!(
!fields.contains_key("integration_latency"),
"integrationLatency should be skipped"
);
assert!(
!fields.contains_key("usage_identifier_key"),
"usageIdentifierKey should be skipped"
);
}
#[test]
fn test_extract_authorizer_v1_non_string_values() {
use aws_lambda_events::apigw::{
ApiGatewayProxyRequestContext, ApiGatewayRequestAuthorizer,
};
let mut authorizer = ApiGatewayRequestAuthorizer::default();
authorizer
.fields
.insert("maxAge".to_string(), serde_json::json!(3600));
authorizer
.fields
.insert("isAdmin".to_string(), serde_json::json!(true));
let mut v1_ctx = ApiGatewayProxyRequestContext::default();
v1_ctx.authorizer = authorizer;
let req =
request_with_context(lambda_http::request::RequestContext::ApiGatewayV1(v1_ctx));
let fields = extract_authorizer_context(&req);
assert_eq!(fields.get("max_age"), Some(&"3600".to_string()));
assert_eq!(fields.get("is_admin"), Some(&"true".to_string()));
}
#[test]
fn test_extract_authorizer_v1_empty() {
use aws_lambda_events::apigw::ApiGatewayProxyRequestContext;
let v1_ctx = ApiGatewayProxyRequestContext::default();
let req =
request_with_context(lambda_http::request::RequestContext::ApiGatewayV1(v1_ctx));
let fields = extract_authorizer_context(&req);
assert!(fields.is_empty());
}
#[test]
fn test_extract_authorizer_v2_lambda_fields() {
use aws_lambda_events::apigw::{
ApiGatewayRequestAuthorizer, ApiGatewayV2httpRequestContext,
};
let mut authorizer = ApiGatewayRequestAuthorizer::default();
authorizer
.fields
.insert("userId".to_string(), serde_json::json!("user-v2"));
authorizer
.fields
.insert("scope".to_string(), serde_json::json!("read write"));
let mut v2_ctx = ApiGatewayV2httpRequestContext::default();
v2_ctx.authorizer = Some(authorizer);
let req =
request_with_context(lambda_http::request::RequestContext::ApiGatewayV2(v2_ctx));
let fields = extract_authorizer_context(&req);
assert_eq!(fields.get("user_id"), Some(&"user-v2".to_string()));
assert_eq!(fields.get("scope"), Some(&"read write".to_string()));
}
}
}