From d841cd0dda96e589a73d53960f73a77302ef29a2 Mon Sep 17 00:00:00 2001 From: "a.tolmachev" Date: Mon, 6 Apr 2026 09:49:22 +0300 Subject: [PATCH] feat: align mcp transport with streamable http --- Cargo.lock | 13 ++ TASKS.md | 12 +- apps/mcp-server/Cargo.toml | 1 + apps/mcp-server/src/app.rs | 254 ++++++++++++++++++++++++++++++++---- apps/mcp-server/src/main.rs | 207 +++++++++++++++++++++++++++++ 5 files changed, 459 insertions(+), 28 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 3410d48..4e60815 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -692,6 +692,17 @@ version = "0.3.32" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cecba35d7ad927e23624b22ad55235f2239cfa44fd10428eecbeba6d6a717718" +[[package]] +name = "futures-macro" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e835b70203e41293343137df5c0664546da5745f82ec9b84d40be8336958447b" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "futures-sink" version = "0.3.32" @@ -712,6 +723,7 @@ checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" dependencies = [ "futures-core", "futures-io", + "futures-macro", "futures-sink", "futures-task", "memchr", @@ -1250,6 +1262,7 @@ dependencies = [ "crank-registry", "crank-runtime", "crank-schema", + "futures-util", "reqwest", "serde", "serde_json", diff --git a/TASKS.md b/TASKS.md index 22b8fbd..9413669 100644 --- a/TASKS.md +++ b/TASKS.md @@ -2,18 +2,20 @@ ## Current -### `feat/streaming-implementation-spec` +### `feat/mcp-streamable-http-alignment` Status: completed DoD: -- Streaming slices are defined file-by-file -- Acceptance criteria and tests are explicit for each slice -- State transitions and sequence outlines are documented +- `mcp-server` transport semantics match MCP Streamable HTTP +- `GET` SSE stream is supported for initialized sessions +- `POST` can return JSON or SSE depending on negotiated response mode +- session and protocol headers are validated consistently +- transport tests cover JSON mode, SSE mode, GET stream and session deletion ## Next -- `feat/mcp-streamable-http-alignment` +- `feat/streaming-core-model` ## Backlog diff --git a/apps/mcp-server/Cargo.toml b/apps/mcp-server/Cargo.toml index 6f419a6..3dfe7f1 100644 --- a/apps/mcp-server/Cargo.toml +++ b/apps/mcp-server/Cargo.toml @@ -12,6 +12,7 @@ crank-core = { path = "../../crates/crank-core" } crank-registry = { path = "../../crates/crank-registry" } crank-runtime = { path = "../../crates/crank-runtime" } crank-schema = { path = "../../crates/crank-schema" } +futures-util = "0.3" serde.workspace = true serde_json.workspace = true sha2.workspace = true diff --git a/apps/mcp-server/src/app.rs b/apps/mcp-server/src/app.rs index 8a980b3..b86e677 100644 --- a/apps/mcp-server/src/app.rs +++ b/apps/mcp-server/src/app.rs @@ -1,4 +1,5 @@ use std::{ + convert::Infallible, sync::Arc, time::{Duration, Instant}, }; @@ -10,7 +11,10 @@ use axum::{ HeaderMap, HeaderValue, StatusCode, header::{self, ACCEPT, AUTHORIZATION}, }, - response::{IntoResponse, Response}, + response::{ + IntoResponse, Response, + sse::{Event, KeepAlive, Sse}, + }, routing::get, }; use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD}; @@ -20,6 +24,7 @@ use crank_core::{ }; use crank_registry::{CreateInvocationLogRequest, PostgresRegistry, PublishedAgentTool}; use crank_runtime::{RuntimeError, RuntimeExecutor, RuntimeOperation}; +use futures_util::stream; use serde::{Deserialize, Serialize}; use serde_json::{Value, json}; use sha2::{Digest, Sha256}; @@ -38,6 +43,12 @@ use crate::{ const HEADER_MCP_SESSION_ID: &str = "MCP-Session-Id"; const HEADER_MCP_PROTOCOL_VERSION: &str = "MCP-Protocol-Version"; +#[derive(Clone, Copy)] +enum ResponseMode { + Json, + Sse, +} + #[derive(Clone)] pub struct AppState { registry: PostgresRegistry, @@ -100,8 +111,53 @@ async fn health() -> Json { })) } -async fn mcp_get(Path(_path): Path) -> Response { - StatusCode::METHOD_NOT_ALLOWED.into_response() +async fn mcp_get( + Path(path): Path, + State(state): State>, + headers: HeaderMap, +) -> Response { + if let Err(status) = validate_origin(&state.allowed_origins, &headers) { + return status.into_response(); + } + + if let Err(status) = validate_get_accept_header(&headers) { + return status.into_response(); + } + + if let Err(status) = + require_platform_api_key(&state, &path, &headers, PlatformApiKeyScope::Read).await + { + return status.into_response(); + } + + let session_id = match session_id_from_headers(&headers) { + Ok(Some(session_id)) => session_id, + Ok(None) => return StatusCode::BAD_REQUEST.into_response(), + Err(status) => return status.into_response(), + }; + + let Some(session) = state.sessions.get(&session_id).await else { + return StatusCode::NOT_FOUND.into_response(); + }; + + if session.workspace_slug != path.workspace_slug || session.agent_slug != path.agent_slug { + return StatusCode::NOT_FOUND.into_response(); + } + + if !session.initialized { + return StatusCode::BAD_REQUEST.into_response(); + } + + if let Err(status) = validate_session_protocol_version(&headers, &session.protocol_version) { + return status.into_response(); + } + + sse_response( + StatusCode::OK, + stream::pending::>(), + Some(&session_id), + Some(&session.protocol_version), + ) } async fn mcp_delete( @@ -145,9 +201,10 @@ async fn mcp_post( return status.into_response(); } - if let Err(status) = validate_accept_header(&headers) { - return status.into_response(); - } + let response_mode = match negotiate_post_response_mode(&headers) { + Ok(mode) => mode, + Err(status) => return status.into_response(), + }; if is_response(&message) || is_notification(&message) && method_name(&message).is_none() { return StatusCode::ACCEPTED.into_response(); @@ -158,6 +215,18 @@ async fn mcp_post( Err(status) => return status.into_response(), }; + if let Some(session_id) = headers.get(HEADER_MCP_SESSION_ID) { + if let Ok(session_id) = session_id.to_str() { + if let Some(session) = state.sessions.get(session_id).await { + if let Err(status) = + validate_session_protocol_version(&headers, &session.protocol_version) + { + return status.into_response(); + } + } + } + } + let required_scope = match method_name(&message) { Some("tools/call") => PlatformApiKeyScope::Write, _ => PlatformApiKeyScope::Read, @@ -168,7 +237,7 @@ async fn mcp_post( match method_name(&message) { Some("initialize") if is_request(&message) => { - handle_initialize(state, &path, &message).await + handle_initialize(state, &path, &message, response_mode).await } Some("notifications/initialized") if is_notification(&message) => { handle_initialized_notification(state, &path, &headers).await @@ -180,14 +249,15 @@ async fn mcp_post( Err(response) => return response, }; - json_response( + transport_response( StatusCode::OK, jsonrpc_result( request_id(&message), json!({ "protocolVersion": session.protocol_version }), ), + response_mode, None, - Some(&protocol_version), + Some(&session.protocol_version), ) } Some("tools/list") if is_request(&message) => { @@ -205,9 +275,10 @@ async fn mcp_post( Ok(tools) => { let definitions = tools.iter().map(tool_definition).collect::>(); - json_response( + transport_response( StatusCode::OK, jsonrpc_result(request_id(&message), json!({ "tools": definitions })), + response_mode, None, Some(&session.protocol_version), ) @@ -224,9 +295,10 @@ async fn mcp_post( let tool_call_params: ToolCallParams = match serde_json::from_value(params(&message)) { Ok(value) => value, Err(error) => { - return json_response( + return transport_response( StatusCode::OK, jsonrpc_error(request_id(&message), -32602, error.to_string()), + response_mode, None, Some(&session.protocol_version), ); @@ -255,7 +327,7 @@ async fn mcp_post( let started_at = Instant::now(); match state.runtime.execute(&runtime_operation, &arguments).await { - Ok(output) => json_response( + Ok(output) => transport_response( StatusCode::OK, { let _ = persist_invocation( @@ -287,6 +359,7 @@ async fn mcp_post( }), ) }, + response_mode, None, Some(&session.protocol_version), ), @@ -306,7 +379,7 @@ async fn mcp_post( }, ) .await; - json_response( + transport_response( StatusCode::OK, jsonrpc_result( request_id(&message), @@ -326,19 +399,21 @@ async fn mcp_post( "isError": true }), ), + response_mode, None, Some(&session.protocol_version), ) } } } - Ok(None) => json_response( + Ok(None) => transport_response( StatusCode::OK, jsonrpc_error( request_id(&message), -32602, format!("tool {} was not found", tool_call_params.name), ), + response_mode, None, Some(&session.protocol_version), ), @@ -349,19 +424,21 @@ async fn mcp_post( let _ = method; StatusCode::ACCEPTED.into_response() } - Some(method) => json_response( + Some(method) => transport_response( StatusCode::OK, jsonrpc_error( request_id(&message), -32601, format!("method {method} is not supported"), ), + response_mode, None, Some(&protocol_version), ), - None => json_response( + None => transport_response( StatusCode::BAD_REQUEST, jsonrpc_error(Value::Null, -32600, "invalid JSON-RPC message"), + response_mode, None, Some(&protocol_version), ), @@ -372,13 +449,15 @@ async fn handle_initialize( state: Arc, path: &AgentRoutePath, message: &Value, + response_mode: ResponseMode, ) -> Response { let initialize_params: InitializeParams = match serde_json::from_value(params(message)) { Ok(value) => value, Err(error) => { - return json_response( + return transport_response( StatusCode::OK, jsonrpc_error(request_id(message), -32602, error.to_string()), + response_mode, None, Some(DEFAULT_PROTOCOL_VERSION), ); @@ -386,7 +465,7 @@ async fn handle_initialize( }; let Some(protocol_version) = negotiated_protocol_version(&initialize_params.protocol_version) else { - return json_response( + return transport_response( StatusCode::OK, jsonrpc_error( request_id(message), @@ -396,6 +475,7 @@ async fn handle_initialize( initialize_params.protocol_version ), ), + response_mode, None, Some(DEFAULT_PROTOCOL_VERSION), ); @@ -405,7 +485,7 @@ async fn handle_initialize( .create(protocol_version, &path.workspace_slug, &path.agent_slug) .await; - json_response( + transport_response( StatusCode::OK, jsonrpc_result( request_id(message), @@ -422,6 +502,7 @@ async fn handle_initialize( } }), ), + response_mode, Some(&session_id), Some(protocol_version), ) @@ -567,16 +648,69 @@ fn validate_origin( Err(StatusCode::FORBIDDEN) } -fn validate_accept_header(headers: &HeaderMap) -> Result<(), StatusCode> { +fn negotiate_post_response_mode(headers: &HeaderMap) -> Result { let Some(accept) = headers.get(ACCEPT) else { return Err(StatusCode::BAD_REQUEST); }; let Ok(accept) = accept.to_str() else { return Err(StatusCode::BAD_REQUEST); }; - let normalized = accept.to_ascii_lowercase(); - if normalized.contains("application/json") && normalized.contains("text/event-stream") { + let mut saw_json = false; + let mut saw_sse = false; + let mut preferred = None; + + for part in accept.split(',') { + let media_type = part + .split(';') + .next() + .unwrap_or_default() + .trim() + .to_ascii_lowercase(); + + match media_type.as_str() { + "application/json" => { + saw_json = true; + if preferred.is_none() { + preferred = Some(ResponseMode::Json); + } + } + "text/event-stream" => { + saw_sse = true; + if preferred.is_none() { + preferred = Some(ResponseMode::Sse); + } + } + _ => {} + } + } + + if saw_json && saw_sse { + return preferred.ok_or(StatusCode::NOT_ACCEPTABLE); + } + + Err(StatusCode::NOT_ACCEPTABLE) +} + +fn validate_get_accept_header(headers: &HeaderMap) -> Result<(), StatusCode> { + let Some(accept) = headers.get(ACCEPT) else { + return Err(StatusCode::BAD_REQUEST); + }; + let Ok(accept) = accept.to_str() else { + return Err(StatusCode::BAD_REQUEST); + }; + + if accept + .split(',') + .map(|part| { + part.split(';') + .next() + .unwrap_or_default() + .trim() + .to_ascii_lowercase() + }) + .any(|media_type| media_type == "text/event-stream") + { return Ok(()); } @@ -598,6 +732,28 @@ fn protocol_version_from_headers(headers: &HeaderMap) -> Result Result<(), StatusCode> { + let Some(version) = headers.get(HEADER_MCP_PROTOCOL_VERSION) else { + return Ok(()); + }; + let Ok(version) = version.to_str() else { + return Err(StatusCode::BAD_REQUEST); + }; + + if negotiated_protocol_version(version).is_none() { + return Err(StatusCode::BAD_REQUEST); + } + + if version != negotiated_session_version { + return Err(StatusCode::BAD_REQUEST); + } + + Ok(()) +} + fn session_id_from_headers(headers: &HeaderMap) -> Result, StatusCode> { let Some(session_id) = headers.get(HEADER_MCP_SESSION_ID) else { return Ok(None); @@ -610,9 +766,10 @@ fn session_id_from_headers(headers: &HeaderMap) -> Result, Status } fn internal_jsonrpc_error(message: &Value, error: impl std::fmt::Display) -> Response { - json_response( + transport_response( StatusCode::INTERNAL_SERVER_ERROR, jsonrpc_error(request_id(message), -32603, error.to_string()), + ResponseMode::Json, None, Some(DEFAULT_PROTOCOL_VERSION), ) @@ -725,6 +882,57 @@ fn json_response( response } +fn transport_response( + status: StatusCode, + payload: Value, + response_mode: ResponseMode, + session_id: Option<&str>, + protocol_version: Option<&str>, +) -> Response { + if status == StatusCode::OK && matches!(response_mode, ResponseMode::Sse) { + let payload = payload.to_string(); + let stream = stream::once(async move { Ok(Event::default().data(payload)) }); + + return sse_response(status, stream, session_id, protocol_version); + } + + json_response(status, payload, session_id, protocol_version) +} + +fn sse_response( + status: StatusCode, + stream: S, + session_id: Option<&str>, + protocol_version: Option<&str>, +) -> Response +where + S: futures_util::stream::Stream> + Send + 'static, +{ + let mut response = ( + status, + Sse::new(stream).keep_alive(KeepAlive::new().interval(Duration::from_secs(15))), + ) + .into_response(); + + if let Some(session_id) = session_id { + response.headers_mut().insert( + HEADER_MCP_SESSION_ID, + HeaderValue::from_str(session_id) + .unwrap_or_else(|_| HeaderValue::from_static("invalid")), + ); + } + + if let Some(protocol_version) = protocol_version { + response.headers_mut().insert( + HEADER_MCP_PROTOCOL_VERSION, + HeaderValue::from_str(protocol_version) + .unwrap_or_else(|_| HeaderValue::from_static(CURRENT_PROTOCOL_VERSION)), + ); + } + + response +} + fn tool_definition(tool: &PublishedAgentTool) -> Value { json!({ "name": tool.tool_name, diff --git a/apps/mcp-server/src/main.rs b/apps/mcp-server/src/main.rs index 12617a3..7fc4fe9 100644 --- a/apps/mcp-server/src/main.rs +++ b/apps/mcp-server/src/main.rs @@ -370,6 +370,213 @@ mod tests { assert_eq!(tools_list["error"]["code"], -32002); } + #[tokio::test] + async fn initialize_can_return_sse_response_when_client_prefers_event_stream() { + let registry = test_registry().await; + publish_agent_with_bindings(®istry, "sales-sse-init", vec![]).await; + let api_key = create_platform_api_key( + ®istry, + "mcp-sse-init", + &[PlatformApiKeyScope::Read, PlatformApiKeyScope::Write], + ) + .await; + let base_url = spawn_mcp_server(build_app( + registry, + Duration::from_millis(0), + Some("https://crank.example.com".to_owned()), + )) + .await; + let client = reqwest::Client::new(); + let response = client + .post(agent_mcp_url(&base_url, "sales-sse-init")) + .header(header::ACCEPT, "text/event-stream, application/json") + .header(header::AUTHORIZATION, format!("Bearer {api_key}")) + .json(&json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "initialize", + "params": { + "protocolVersion": "2025-11-25" + } + })) + .send() + .await + .unwrap(); + + assert_eq!(response.status(), reqwest::StatusCode::OK); + assert_eq!( + response + .headers() + .get(header::CONTENT_TYPE) + .unwrap() + .to_str() + .unwrap(), + "text/event-stream" + ); + assert!(response.headers().get("MCP-Session-Id").is_some()); + + let body = response.text().await.unwrap(); + assert!(body.contains("\"jsonrpc\":\"2.0\"")); + assert!(body.contains("\"protocolVersion\":\"2025-11-25\"")); + } + + #[tokio::test] + async fn get_opens_sse_stream_for_initialized_session() { + let registry = test_registry().await; + publish_agent_with_bindings(®istry, "sales-get-sse", vec![]).await; + let api_key = create_platform_api_key( + ®istry, + "mcp-get-sse", + &[PlatformApiKeyScope::Read, PlatformApiKeyScope::Write], + ) + .await; + let base_url = spawn_mcp_server(build_app( + registry, + Duration::from_millis(0), + Some("https://crank.example.com".to_owned()), + )) + .await; + let client = reqwest::Client::new(); + let mcp_url = agent_mcp_url(&base_url, "sales-get-sse"); + let initialized_session = initialize_session(&client, &mcp_url, &api_key).await; + + let response = client + .get(&mcp_url) + .header(header::ACCEPT, "text/event-stream") + .header(header::AUTHORIZATION, format!("Bearer {api_key}")) + .header("MCP-Session-Id", &initialized_session) + .header("MCP-Protocol-Version", "2025-11-25") + .send() + .await + .unwrap(); + + assert_eq!(response.status(), reqwest::StatusCode::OK); + assert_eq!( + response + .headers() + .get(header::CONTENT_TYPE) + .unwrap() + .to_str() + .unwrap(), + "text/event-stream" + ); + assert_eq!( + response + .headers() + .get("MCP-Session-Id") + .unwrap() + .to_str() + .unwrap(), + initialized_session + ); + } + + #[tokio::test] + async fn get_requires_session_header() { + let registry = test_registry().await; + publish_agent_with_bindings(®istry, "sales-get-sse-missing", vec![]).await; + let api_key = create_platform_api_key( + ®istry, + "mcp-get-sse-missing", + &[PlatformApiKeyScope::Read], + ) + .await; + let base_url = spawn_mcp_server(build_app( + registry, + Duration::from_millis(0), + Some("https://crank.example.com".to_owned()), + )) + .await; + let client = reqwest::Client::new(); + + let response = client + .get(agent_mcp_url(&base_url, "sales-get-sse-missing")) + .header(header::ACCEPT, "text/event-stream") + .header(header::AUTHORIZATION, format!("Bearer {api_key}")) + .send() + .await + .unwrap(); + + assert_eq!(response.status(), reqwest::StatusCode::BAD_REQUEST); + } + + #[tokio::test] + async fn delete_terminates_transport_session() { + let registry = test_registry().await; + publish_agent_with_bindings(®istry, "sales-delete-session", vec![]).await; + let api_key = create_platform_api_key( + ®istry, + "mcp-delete-session", + &[PlatformApiKeyScope::Read], + ) + .await; + let base_url = spawn_mcp_server(build_app( + registry, + Duration::from_millis(0), + Some("https://crank.example.com".to_owned()), + )) + .await; + let client = reqwest::Client::new(); + let mcp_url = agent_mcp_url(&base_url, "sales-delete-session"); + let initialized_session = initialize_session(&client, &mcp_url, &api_key).await; + + let delete_response = client + .delete(&mcp_url) + .header(header::AUTHORIZATION, format!("Bearer {api_key}")) + .header("MCP-Session-Id", &initialized_session) + .header("MCP-Protocol-Version", "2025-11-25") + .send() + .await + .unwrap(); + + assert_eq!(delete_response.status(), reqwest::StatusCode::NO_CONTENT); + + let after_delete = client + .get(&mcp_url) + .header(header::ACCEPT, "text/event-stream") + .header(header::AUTHORIZATION, format!("Bearer {api_key}")) + .header("MCP-Session-Id", &initialized_session) + .header("MCP-Protocol-Version", "2025-11-25") + .send() + .await + .unwrap(); + + assert_eq!(after_delete.status(), reqwest::StatusCode::NOT_FOUND); + } + + #[tokio::test] + async fn rejects_get_with_protocol_version_mismatch() { + let registry = test_registry().await; + publish_agent_with_bindings(®istry, "sales-get-bad-version", vec![]).await; + let api_key = create_platform_api_key( + ®istry, + "mcp-get-bad-version", + &[PlatformApiKeyScope::Read], + ) + .await; + let base_url = spawn_mcp_server(build_app( + registry, + Duration::from_millis(0), + Some("https://crank.example.com".to_owned()), + )) + .await; + let client = reqwest::Client::new(); + let mcp_url = agent_mcp_url(&base_url, "sales-get-bad-version"); + let initialized_session = initialize_session(&client, &mcp_url, &api_key).await; + + let response = client + .get(&mcp_url) + .header(header::ACCEPT, "text/event-stream") + .header(header::AUTHORIZATION, format!("Bearer {api_key}")) + .header("MCP-Session-Id", &initialized_session) + .header("MCP-Protocol-Version", "2025-06-18") + .send() + .await + .unwrap(); + + assert_eq!(response.status(), reqwest::StatusCode::BAD_REQUEST); + } + #[tokio::test] async fn refreshes_published_tools_without_restart() { let registry = test_registry().await;