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 331dffaf98.

* dont use new path

* cleaning

* cleaner way of closing sessions
This commit is contained in:
centdix
2025-06-12 10:49:22 +02:00
committed by GitHub
parent 009bfebfd1
commit bef5ed8c24
4 changed files with 152 additions and 71 deletions

30
backend/Cargo.lock generated
View File

@@ -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"

View File

@@ -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

View File

@@ -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::<Arc<LocalSessionManager>>::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::<OptAuthed>())
.layer(cors.clone()),
)
.nest("/mcp/w/:workspace_id", mcp_router)
.nest("/mcp/w/:workspace_id/sse", mcp_router)
.layer(from_extractor::<OptAuthed>())
.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");
}
});

View File

@@ -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::<ApiAuthed>()
.ok_or_else(|| Error::internal_error("ApiAuthed not found", None))?;
let db = context
.req_extensions
.get::<DB>()
.ok_or_else(|| Error::internal_error("DB not found", None))?;
let user_db = context
.req_extensions
.get::<UserDB>()
.ok_or_else(|| Error::internal_error("UserDB not found", None))?;
let http_parts = context
.extensions
.get::<axum::http::request::Parts>()
.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::<ApiAuthed>().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::<DB>().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::<UserDB>().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::<WorkspaceId>()
.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<PaginatedRequestParam>,
mut _context: RequestContext<RoleServer>,
) -> Result<ListToolsResult, Error> {
let workspace_id = _context.workspace_id.clone();
let db = _context
.req_extensions
.get::<DB>()
.ok_or_else(|| Error::internal_error("DB not found", None))?;
let user_db = _context
.req_extensions
.get::<UserDB>()
.ok_or_else(|| Error::internal_error("UserDB not found", None))?;
let authed = _context
.req_extensions
.get::<ApiAuthed>()
.ok_or_else(|| Error::internal_error("ApiAuthed not found", None))?;
let http_parts = _context
.extensions
.get::<axum::http::request::Parts>()
.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::<DB>().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::<UserDB>().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::<ApiAuthed>().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::<WorkspaceId>()
.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<String>,
mut request: Request<axum::body::Body>,
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<LocalSessionManager>)> {
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<LocalSessionManager>) {
let session_ids_to_close = {
let sessions_map = session_manager.sessions.read().await;
sessions_map.keys().cloned().collect::<Vec<_>>()
};
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::<Vec<_>>();
futures::future::join_all(close_futures).await;
}
}