draft for http streamable usage

This commit is contained in:
centdix
2025-06-10 14:54:39 +02:00
parent ae81b4f456
commit 5361b0ff1f
4 changed files with 115 additions and 53 deletions
+25 -5
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"
+1 -1
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
+20 -20
View File
@@ -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?;
+69 -27
View File
@@ -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::<ApiAuthed>()
.ok_or_else(|| Error::internal_error("ApiAuthed not found", None))?;
let db = context
.req_extensions
.extensions
.get::<DB>()
.ok_or_else(|| Error::internal_error("DB not found", None))?;
let user_db = context
.req_extensions
.extensions
.get::<UserDB>()
.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<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))?;
println!("list_tools");
println!(
"[HANDLER DEBUG] _context.extensions contains http::Parts? {:?}",
_context.extensions.get::<http::request::Parts>().is_some()
);
// 1. Get http::request::Parts from MCP extensions
let http_parts = _context
.extensions
.get::<axum::http::request::Parts>() // 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::<axum::extract::Extension<MyTestMarker>>() // 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::<axum::extract::Extension<OptAuthed>>() // 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::<DB>().ok_or_else(|| {
println!("DB not found");
Error::internal_error("DB not found", None)
})?;
let user_db = _context
.req_extensions
.extensions
.get::<UserDB>()
.ok_or_else(|| Error::internal_error("UserDB not found", None))?;
let authed = _context
.req_extensions
.extensions
.get::<ApiAuthed>()
.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<Router> {
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)
}