mirror of
https://github.com/windmill-labs/windmill.git
synced 2026-08-23 16:00:38 +00:00
draft for http streamable usage
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
|
||||
|
||||
@@ -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?;
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user