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..d1aacc967c 100644 --- a/backend/windmill-api/src/lib.rs +++ b/backend/windmill-api/src/lib.rs @@ -520,17 +520,17 @@ pub async fn run_server( // Setup MCP server #[allow(unused_variables)] - let (mcp_router, mcp_main_ct, mcp_service_ct) = { + let mcp_router = { #[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 = setup_mcp_server(addr).await?; + // #[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 } else { - (Router::new(), None, None) + Router::new() } #[cfg(not(feature = "mcp"))] @@ -818,18 +818,18 @@ 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."); - } - } + // #[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."); + // } + // } }); server.await?; diff --git a/backend/windmill-api/src/mcp.rs b/backend/windmill-api/src/mcp.rs index 431e121ba8..cc6ed09e21 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 std::{borrow::Cow, time::Duration}; use axum::body::to_bytes; -use axum::Router; -use rmcp::transport::sse_server::{SseServer, SseServerConfig}; +use axum::{Extension, Router}; use rmcp::{ handler::server::ServerHandler, model::*, @@ -24,13 +23,20 @@ use windmill_common::{DB, HUB_BASE_URL}; use windmill_common::scripts::{get_full_hub_script_by_path, Schema}; -use crate::db::ApiAuthed; use crate::jobs::{ run_wait_result_flow_by_path_internal, run_wait_result_script_by_path_internal, RunJobQuery, }; use crate::HTTP_CLIENT; +use crate::{db::ApiAuthed, users::OptAuthed}; +use rmcp::transport::streamable_http_server::{ + session::local::LocalSessionManager, StreamableHttpService, +}; +use rmcp::transport::StreamableHttpServerConfig; use windmill_common::utils::{query_elems_from_hub, StripPath}; +#[derive(Debug, Clone)] +pub struct MyTestMarker(pub String); + /// Transforms the path for workspace scripts/flows. /// /// This function takes a path and a type string. @@ -857,15 +863,15 @@ impl ServerHandler for Runner { }; let authed = context - .req_extensions + .extensions .get::() .ok_or_else(|| Error::internal_error("ApiAuthed not found", None))?; let db = context - .req_extensions + .extensions .get::() .ok_or_else(|| Error::internal_error("DB not found", None))?; let user_db = context - .req_extensions + .extensions .get::() .ok_or_else(|| Error::internal_error("UserDB not found", None))?; let args = parse_args(request.arguments)?; @@ -873,11 +879,11 @@ impl ServerHandler for Runner { let (tool_type, path, is_hub) = Runner::reverse_transform(&request.name).unwrap_or_default(); + let w_id = "admin".to_string(); 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, &w_id, &tool_type).await? }; let schema_obj = if let Some(ref s) = item_schema { @@ -907,7 +913,7 @@ impl ServerHandler for Runner { windmill_queue::PushArgsOwned::default() }; - let w_id = context.workspace_id.clone(); + let w_id = "admin".to_string(); let script_or_flow_path = if is_hub { StripPath(format!("hub/{}", path)) } else { @@ -978,17 +984,53 @@ 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))?; + println!("list_tools"); + + println!( + "[HANDLER DEBUG] _context.extensions contains http::Parts? {:?}", + _context.extensions.get::().is_some() + ); + + // 1. Get http::request::Parts from MCP extensions + let http_parts = _context + .extensions + .get::() // Ensure http::request::Parts is in scope + .ok_or_else(|| { + println!("CRITICAL: http::request::Parts not found in MCP extensions. Was it injected by the transport layer?"); + Error::internal_error("http::request::Parts not found", None) + })?; + + println!("http_parts: {:?}", http_parts); + + // 2. Get your Axum extension (e.g., DB) from http_parts.extensions + // Axum stores extensions wrapped in axum::extract::Extension + let db_axum_extension = http_parts + .extensions + .get::>() // Ensure axum::extract::Extension is in scope + .ok_or_else(|| { + println!("DB Axum extension not found"); + // Error::internal_error("DB Axum extension not found", None) + }); + + let authed_axum_extension = http_parts + .extensions + .get::>() // Ensure axum::extract::Extension is in scope + .ok_or_else(|| { + println!("Authed Axum extension not found"); + Error::internal_error("Authed Axum extension not found", None) + })?; + + let workspace_id = "admin".to_string(); + let db = _context.extensions.get::().ok_or_else(|| { + println!("DB not found"); + Error::internal_error("DB not found", None) + })?; let user_db = _context - .req_extensions + .extensions .get::() .ok_or_else(|| Error::internal_error("UserDB not found", None))?; let authed = _context - .req_extensions + .extensions .get::() .ok_or_else(|| Error::internal_error("ApiAuthed not found", None))?; let owned_scope = authed.scopes.as_ref().and_then(|scopes| { @@ -1079,6 +1121,8 @@ impl ServerHandler for Runner { ); } + println!("tools: {:?}", tools); + Ok(ListToolsResult { tools, next_cursor: None }) } @@ -1127,15 +1171,13 @@ 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, - }; +pub async fn setup_mcp_server(addr: SocketAddr) -> anyhow::Result { + let config = StreamableHttpServerConfig { sse_keep_alive: None, stateful_mode: true }; + let service = + StreamableHttpService::new(Runner::new, LocalSessionManager::default().into(), config); - Ok(SseServer::new(config)) + let router = axum::Router::new() + .nest_service("/", service) + .layer(Extension(MyTestMarker("hello_from_axum_layer".to_string()))); + Ok(router) }