mirror of
https://github.com/windmill-labs/windmill.git
synced 2026-08-20 08:01:35 +00:00
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:
Generated
+25
-5
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
});
|
||||
|
||||
|
||||
+112
-43
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user