entity_graph_mcp/
transport.rs

1//! Transport helpers.
2//!
3//! * [`run_stdio`]   — attach to stdin/stdout for local MCP hosts.
4//! * [`run_http`]    — bind an Axum HTTP server with `StreamableHttpService`.
5
6use std::{net::SocketAddr, sync::Arc};
7
8use anyhow::Result;
9use axum::{extract::State, routing::get, Router};
10use rmcp::{
11    serve_server,
12    transport::{
13        streamable_http_server::session::local::LocalSessionManager, StreamableHttpServerConfig,
14        StreamableHttpService,
15    },
16};
17use tokio::net::TcpListener;
18use tokio_util::sync::CancellationToken;
19use tower_http::{cors::CorsLayer, trace::TraceLayer};
20use tracing::info;
21
22use crate::server::EntityGraphServer;
23
24// ── stdio transport ───────────────────────────────────────────────────────────
25
26/// Run the MCP server on stdin / stdout.
27///
28/// Blocks until the client disconnects (EOF on stdin).
29pub async fn run_stdio(server: EntityGraphServer) -> Result<()> {
30    info!("entity-graph-mcp: starting stdio transport");
31
32    let (stdin, stdout) = rmcp::transport::stdio();
33    let service = serve_server(server, (stdin, stdout)).await?;
34    service.waiting().await?;
35
36    Ok(())
37}
38
39// ── HTTP transport ────────────────────────────────────────────────────────────
40
41/// Shared state available to HTTP health/info routes.
42#[derive(Clone)]
43struct AppState {
44    server_name: String,
45    server_version: String,
46}
47
48/// Run the MCP server over Streamable HTTP using Axum.
49///
50/// The MCP endpoint is mounted at `/mcp` (POST + GET, as required by the
51/// Streamable HTTP spec). A lightweight `/health` route is also provided.
52///
53/// # Arguments
54///
55/// * `server` — the `EntityGraphServer` instance (will be cloned per session).
56/// * `addr`   — the socket address to bind (e.g. `"0.0.0.0:8080".parse()?`).
57/// * `cancel` — cancellation token; call `.cancel()` to initiate shutdown.
58pub async fn run_http(
59    server: EntityGraphServer,
60    addr: SocketAddr,
61    cancel: CancellationToken,
62) -> Result<()> {
63    info!("entity-graph-mcp: starting Streamable HTTP transport on {addr}");
64
65    let session_manager = Arc::new(LocalSessionManager::default());
66
67    // `StreamableHttpServerConfig` is `#[non_exhaustive]`; use its builder
68    // methods instead of struct literal syntax.
69    let config = StreamableHttpServerConfig::default()
70        .with_stateful_mode(true)
71        .with_cancellation_token(cancel.clone())
72        .with_allowed_hosts(["localhost", "127.0.0.1", "0.0.0.0"]);
73
74    // Clone the server for each new MCP session.
75    let server_factory = {
76        let server = server.clone();
77        move || Ok(server.clone())
78    };
79
80    let mcp_service = StreamableHttpService::new(server_factory, session_manager, config);
81
82    let state = AppState {
83        server_name: env!("CARGO_PKG_NAME").to_owned(),
84        server_version: env!("CARGO_PKG_VERSION").to_owned(),
85    };
86
87    let app = Router::new()
88        .route("/health", get(health_handler))
89        .route("/mcp", axum::routing::any_service(mcp_service))
90        .with_state(state)
91        .layer(TraceLayer::new_for_http())
92        .layer(CorsLayer::permissive());
93
94    let listener = TcpListener::bind(addr).await?;
95    info!("entity-graph-mcp: HTTP server listening on {addr}");
96
97    axum::serve(listener, app)
98        .with_graceful_shutdown(async move {
99            cancel.cancelled().await;
100            info!("entity-graph-mcp: HTTP server shutting down");
101        })
102        .await?;
103
104    Ok(())
105}
106
107/// `GET /health` — liveness probe.
108async fn health_handler(State(state): State<AppState>) -> axum::Json<serde_json::Value> {
109    axum::Json(serde_json::json!({
110        "status": "ok",
111        "server": state.server_name,
112        "version": state.server_version,
113    }))
114}