feat: align mcp transport with streamable http

This commit is contained in:
a.tolmachev
2026-04-06 09:49:22 +03:00
parent 633af39c82
commit d841cd0dda
5 changed files with 459 additions and 28 deletions
Generated
+13
View File
@@ -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",
+7 -5
View File
@@ -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
+1
View File
@@ -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
+231 -23
View File
@@ -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<Value> {
}))
}
async fn mcp_get(Path(_path): Path<AgentRoutePath>) -> Response {
StatusCode::METHOD_NOT_ALLOWED.into_response()
async fn mcp_get(
Path(path): Path<AgentRoutePath>,
State(state): State<Arc<AppState>>,
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::<Result<Event, Infallible>>(),
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::<Vec<_>>();
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<AppState>,
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<ResponseMode, 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);
};
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<String, StatusCo
Err(StatusCode::BAD_REQUEST)
}
fn validate_session_protocol_version(
headers: &HeaderMap,
negotiated_session_version: &str,
) -> 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<Option<String>, 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<Option<String>, 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<S>(
status: StatusCode,
stream: S,
session_id: Option<&str>,
protocol_version: Option<&str>,
) -> Response
where
S: futures_util::stream::Stream<Item = Result<Event, Infallible>> + 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,
+207
View File
@@ -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(&registry, "sales-sse-init", vec![]).await;
let api_key = create_platform_api_key(
&registry,
"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(&registry, "sales-get-sse", vec![]).await;
let api_key = create_platform_api_key(
&registry,
"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(&registry, "sales-get-sse-missing", vec![]).await;
let api_key = create_platform_api_key(
&registry,
"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(&registry, "sales-delete-session", vec![]).await;
let api_key = create_platform_api_key(
&registry,
"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(&registry, "sales-get-bad-version", vec![]).await;
let api_key = create_platform_api_key(
&registry,
"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;