From bef5ed8c2486cce59d4efbbbe5cdb3fbfbc5881c Mon Sep 17 00:00:00 2001 From: centdix <40307056+centdix@users.noreply.github.com> Date: Thu, 12 Jun 2025 10:49:22 +0200 Subject: [PATCH] feat(backend): use streamable http in favor of sse for MCP (#5910) * draft for http streamable usage * good stuff * add workspace_id to extensions * fix shutdown * cleaning * fix * adapt frontend * Revert "adapt frontend" This reverts commit 331dffaf98004723ce68270002b7c9ed6690bfec. * dont use new path * cleaning * cleaner way of closing sessions --- backend/Cargo.lock | 30 +++++-- backend/windmill-api/Cargo.toml | 2 +- backend/windmill-api/src/lib.rs | 36 +++----- backend/windmill-api/src/mcp.rs | 155 +++++++++++++++++++++++--------- 4 files changed, 152 insertions(+), 71 deletions(-) diff --git a/backend/Cargo.lock b/backend/Cargo.lock index edaea2c45a..37d481c505 100644 --- a/backend/Cargo.lock +++ b/backend/Cargo.lock @@ -10402,13 +10402,15 @@ dependencies = [ [[package]] name = "rmcp" version = "0.1.5" -source = "git+https://github.com/windmill-labs/rust-sdk#9142b40202e49ca0b6530fa49a3abc8bd1b2fcf0" +source = "git+https://github.com/modelcontextprotocol/rust-sdk#db03f63e76b5b32f65d34a1bd08ae56dab595f60" dependencies = [ - "async-stream", - "axum", - "base64 0.21.7", + "base64 0.22.1", + "bytes", "chrono", "futures", + "http 1.3.1", + "http-body 1.0.1", + "http-body-util", "paste", "pin-project-lite", "rand 0.9.0", @@ -10416,20 +10418,24 @@ dependencies = [ "schemars", "serde", "serde_json", + "sse-stream", "thiserror 2.0.12", "tokio", "tokio-stream", "tokio-util", + "tower-service", "tracing", + "uuid", ] [[package]] name = "rmcp-macros" version = "0.1.5" -source = "git+https://github.com/windmill-labs/rust-sdk#9142b40202e49ca0b6530fa49a3abc8bd1b2fcf0" +source = "git+https://github.com/modelcontextprotocol/rust-sdk#db03f63e76b5b32f65d34a1bd08ae56dab595f60" dependencies = [ "proc-macro2", "quote", + "serde_json", "syn 2.0.101", ] @@ -10967,6 +10973,7 @@ version = "0.8.22" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3fbf2ae1b8bc8e02df939598064d22402220cd5bbcca1c76f7d6a310974d5615" dependencies = [ + "chrono", "dyn-clone", "schemars_derive", "serde", @@ -11880,6 +11887,19 @@ dependencies = [ "uuid", ] +[[package]] +name = "sse-stream" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f649a9f9e91db2ed32f3724516eac2bc09fab77fc33be8f670f5619b9dc6c3f" +dependencies = [ + "bytes", + "futures-util", + "http-body 1.0.1", + "http-body-util", + "pin-project-lite", +] + [[package]] name = "stable_deref_trait" version = "1.2.0" diff --git a/backend/windmill-api/Cargo.toml b/backend/windmill-api/Cargo.toml index cfb175cb0f..8017436a4a 100644 --- a/backend/windmill-api/Cargo.toml +++ b/backend/windmill-api/Cargo.toml @@ -40,7 +40,7 @@ mcp = ["dep:rmcp"] python = [] [dependencies] -rmcp = { git = "https://github.com/windmill-labs/rust-sdk", features = ["transport-sse-server"], optional = true } +rmcp = { git = "https://github.com/modelcontextprotocol/rust-sdk", features=["transport-streamable-http-server", "transport-streamable-http-server-session", "transport-worker"], optional = true } windmill-queue.workspace = true windmill-common = { workspace = true, default-features = false } windmill-audit.workspace = true diff --git a/backend/windmill-api/src/lib.rs b/backend/windmill-api/src/lib.rs index 8a42902508..3eb1b9d183 100644 --- a/backend/windmill-api/src/lib.rs +++ b/backend/windmill-api/src/lib.rs @@ -19,14 +19,16 @@ use crate::oauth2_oss::SlackVerifier; use crate::smtp_server_oss::SmtpServer; #[cfg(feature = "mcp")] -use crate::mcp::{setup_mcp_server, Runner as McpRunner}; +use crate::mcp::{extract_and_store_workspace_id, setup_mcp_server, shutdown_mcp_server}; +#[cfg(feature = "mcp")] +use rmcp::transport::streamable_http_server::session::local::LocalSessionManager; + use crate::tracing_init::MyOnFailure; use crate::{ tracing_init::{MyMakeSpan, MyOnResponse}, users::OptAuthed, webhook_util::WebhookShared, }; - #[cfg(feature = "agent_worker_server")] use agent_workers_oss::AgentCache; @@ -520,21 +522,18 @@ pub async fn run_server( // Setup MCP server #[allow(unused_variables)] - let (mcp_router, mcp_main_ct, mcp_service_ct) = { + let (mcp_router, mcp_session_manager) = { #[cfg(feature = "mcp")] if server_mode || mcp_mode { - let (mcp_sse_server, mcp_router) = setup_mcp_server(addr, "/api/mcp/w/:workspace_id")?; - #[cfg(feature = "mcp")] - let mcp_main_ct = mcp_sse_server.config.ct.clone(); // Token to signal shutdown *to* MCP - #[cfg(feature = "mcp")] - let mcp_service_ct = mcp_sse_server.with_service(McpRunner::new); // Token to wait for MCP *service* shutdown - (mcp_router, Some(mcp_main_ct), Some(mcp_service_ct)) + let (mcp_router, mcp_session_manager) = setup_mcp_server().await?; + let mcp_middleware = axum::middleware::from_fn(extract_and_store_workspace_id); + (mcp_router.layer(mcp_middleware), Some(mcp_session_manager)) } else { - (Router::new(), None, None) + (Router::new(), Option::>::None) } #[cfg(not(feature = "mcp"))] - (Router::new(), None::<()>, None::<()>) + (Router::new(), Option::<()>::None) }; #[cfg(feature = "agent_worker_server")] @@ -660,7 +659,7 @@ pub async fn run_server( .layer(from_extractor::()) .layer(cors.clone()), ) - .nest("/mcp/w/:workspace_id", mcp_router) + .nest("/mcp/w/:workspace_id/sse", mcp_router) .layer(from_extractor::()) .nest("/agent_workers", { #[cfg(feature = "agent_worker_server")] @@ -819,16 +818,9 @@ pub async fn run_server( tracing::info!("Graceful shutdown of server"); #[cfg(feature = "mcp")] - { - if let Some(mcp_main_ct) = mcp_main_ct { - tracing::info!("Received shutdown signal, cancelling MCP server..."); - mcp_main_ct.cancel(); - } - if let Some(mcp_service_ct) = mcp_service_ct { - tracing::info!("Waiting for MCP service cancellation..."); - mcp_service_ct.cancelled().await; - tracing::info!("MCP service cancelled."); - } + if let Some(mcp_session_manager) = mcp_session_manager { + shutdown_mcp_server(mcp_session_manager).await; + tracing::info!("MCP server shutdown"); } }); diff --git a/backend/windmill-api/src/mcp.rs b/backend/windmill-api/src/mcp.rs index 431e121ba8..dcd6fa4ca8 100644 --- a/backend/windmill-api/src/mcp.rs +++ b/backend/windmill-api/src/mcp.rs @@ -1,11 +1,10 @@ use std::borrow::Cow; use std::collections::HashMap; -use std::net::SocketAddr; use std::sync::Arc; use axum::body::to_bytes; use axum::Router; -use rmcp::transport::sse_server::{SseServer, SseServerConfig}; +use axum::{extract::Path, http::Request, middleware::Next, response::Response}; use rmcp::{ handler::server::ServerHandler, model::*, @@ -17,7 +16,6 @@ use serde_json::Value; use sql_builder::prelude::*; use sqlx::FromRow; use tokio::try_join; -use tokio_util::sync::CancellationToken; use windmill_common::db::UserDB; use windmill_common::worker::to_raw_value; use windmill_common::{DB, HUB_BASE_URL}; @@ -29,6 +27,9 @@ use crate::jobs::{ run_wait_result_flow_by_path_internal, run_wait_result_script_by_path_internal, RunJobQuery, }; use crate::HTTP_CLIENT; +use rmcp::transport::streamable_http_server::{ + session::local::LocalSessionManager, SessionManager, StreamableHttpService, +}; use windmill_common::utils::{query_elems_from_hub, StripPath}; /// Transforms the path for workspace scripts/flows. @@ -856,28 +857,44 @@ impl ServerHandler for Runner { }) }; - let authed = context - .req_extensions - .get::() - .ok_or_else(|| Error::internal_error("ApiAuthed not found", None))?; - let db = context - .req_extensions - .get::() - .ok_or_else(|| Error::internal_error("DB not found", None))?; - let user_db = context - .req_extensions - .get::() - .ok_or_else(|| Error::internal_error("UserDB not found", None))?; + let http_parts = context + .extensions + .get::() + .ok_or_else(|| { + tracing::error!("http::request::Parts not found"); + Error::internal_error("http::request::Parts not found", None) + })?; + + let authed = http_parts.extensions.get::().ok_or_else(|| { + tracing::error!("ApiAuthed Axum extension not found"); + Error::internal_error("ApiAuthed Axum extension not found", None) + })?; + let db = http_parts.extensions.get::().ok_or_else(|| { + tracing::error!("DB Axum extension not found"); + Error::internal_error("DB Axum extension not found", None) + })?; + let user_db = http_parts.extensions.get::().ok_or_else(|| { + tracing::error!("UserDB Axum extension not found"); + Error::internal_error("UserDB Axum extension not found", None) + })?; let args = parse_args(request.arguments)?; + let workspace_id = http_parts + .extensions + .get::() + .ok_or_else(|| { + tracing::error!("WorkspaceId not found"); + Error::internal_error("WorkspaceId not found", None) + }) + .map(|w_id| w_id.0.clone())?; + let (tool_type, path, is_hub) = Runner::reverse_transform(&request.name).unwrap_or_default(); let item_schema = if is_hub { Runner::get_hub_script_schema(&format!("hub/{}", path), db).await? } else { - Runner::get_item_schema(&path, user_db, authed, &context.workspace_id, &tool_type) - .await? + Runner::get_item_schema(&path, user_db, authed, &workspace_id, &tool_type).await? }; let schema_obj = if let Some(ref s) = item_schema { @@ -906,8 +923,6 @@ impl ServerHandler for Runner { } else { windmill_queue::PushArgsOwned::default() }; - - let w_id = context.workspace_id.clone(); let script_or_flow_path = if is_hub { StripPath(format!("hub/{}", path)) } else { @@ -922,7 +937,7 @@ impl ServerHandler for Runner { script_or_flow_path, authed.clone(), user_db.clone(), - w_id.clone(), + workspace_id.clone(), push_args, ) .await @@ -934,7 +949,7 @@ impl ServerHandler for Runner { authed.clone(), user_db.clone(), push_args, - w_id.clone(), + workspace_id.clone(), ) .await }; @@ -978,19 +993,38 @@ impl ServerHandler for Runner { _request: Option, mut _context: RequestContext, ) -> Result { - let workspace_id = _context.workspace_id.clone(); - let db = _context - .req_extensions - .get::() - .ok_or_else(|| Error::internal_error("DB not found", None))?; - let user_db = _context - .req_extensions - .get::() - .ok_or_else(|| Error::internal_error("UserDB not found", None))?; - let authed = _context - .req_extensions - .get::() - .ok_or_else(|| Error::internal_error("ApiAuthed not found", None))?; + let http_parts = _context + .extensions + .get::() + .ok_or_else(|| { + tracing::error!("http::request::Parts not found"); + Error::internal_error("http::request::Parts not found", None) + })?; + + let db = http_parts.extensions.get::().ok_or_else(|| { + tracing::error!("DB Axum extension not found"); + Error::internal_error("DB Axum extension not found", None) + })?; + + let user_db = http_parts.extensions.get::().ok_or_else(|| { + tracing::error!("UserDB Axum extension not found"); + Error::internal_error("UserDB Axum extension not found", None) + })?; + + let authed = http_parts.extensions.get::().ok_or_else(|| { + tracing::error!("ApiAuthed Axum extension not found"); + Error::internal_error("ApiAuthed Axum extension not found", None) + })?; + + let workspace_id = http_parts + .extensions + .get::() + .ok_or_else(|| { + tracing::error!("WorkspaceId not found"); + Error::internal_error("WorkspaceId not found", None) + }) + .map(|w_id| w_id.0.clone())?; + let owned_scope = authed.scopes.as_ref().and_then(|scopes| { scopes .iter() @@ -1127,15 +1161,50 @@ impl ServerHandler for Runner { } } -pub fn setup_mcp_server(addr: SocketAddr, path: &str) -> anyhow::Result<(SseServer, Router)> { - let config = SseServerConfig { - bind: addr, - sse_path: "/sse".to_string(), - post_path: "/message".to_string(), - full_message_path: path.to_string(), - ct: CancellationToken::new(), - sse_keep_alive: None, +#[derive(Clone, Debug)] +pub struct WorkspaceId(pub String); + +pub async fn extract_and_store_workspace_id( + Path(params): Path, + mut request: Request, + next: Next, +) -> Response { + let workspace_id = params; + request.extensions_mut().insert(WorkspaceId(workspace_id)); + next.run(request).await +} + +pub async fn setup_mcp_server() -> anyhow::Result<(Router, Arc)> { + let session_manager = Arc::new(LocalSessionManager::default()); + let service_config = Default::default(); + let service = StreamableHttpService::new(Runner::new, session_manager.clone(), service_config); + + let router = axum::Router::new().nest_service("/", service); + Ok((router, session_manager)) +} + +pub async fn shutdown_mcp_server(session_manager: Arc) { + let session_ids_to_close = { + let sessions_map = session_manager.sessions.read().await; + sessions_map.keys().cloned().collect::>() }; - Ok(SseServer::new(config)) + if !session_ids_to_close.is_empty() { + tracing::info!( + "Closing {} active MCP session(s)...", + session_ids_to_close.len() + ); + let close_futures = session_ids_to_close + .iter() + .map(|session_id| { + let manager_clone = session_manager.clone(); + async move { + if let Err(_) = manager_clone.close_session(session_id).await { + tracing::warn!("Error closing MCP session"); + } + } + }) + .collect::>(); + futures::future::join_all(close_futures).await; + } }