Skip to main content

mcp_airlock/
lib.rs

1//! # mcp-airlock
2//!
3//! A local **stdio ⇄ Streamable HTTP bridge** for the [Model Context
4//! Protocol](https://modelcontextprotocol.io). It lets AI clients that launch
5//! local MCP servers talk to remote, OAuth-protected MCP servers, with FAPI 2.0
6//! security: Pushed Authorization Requests, PKCE and DPoP-bound tokens.
7//!
8//! Most people use the `mcp-airlock` binary (see the
9//! [README](https://github.com/ffalcinelli/mcp-airlock#readme) and the
10//! [setup guide](https://github.com/ffalcinelli/mcp-airlock/blob/main/GUIDE.md)).
11//! This crate exposes the same machinery as a library.
12//!
13//! ## How it works
14//!
15//! 1. Each JSON-RPC message read from stdin is POSTed to the MCP server, with
16//!    the headers of its protocol era. Modern (2026-07-28) messages carry
17//!    `MCP-Protocol-Version`, `Mcp-Method`, `Mcp-Name` and `Mcp-Param-*`;
18//!    legacy (`initialize`-based) ones carry the negotiated version and the
19//!    session id.
20//! 2. A `401` or `403 insufficient_scope` activates the *airlock*: requests
21//!    pause while the token is refreshed or, if needed, a browser login runs:
22//!    RFC 9728 → RFC 8414 discovery, PAR with PKCE and `dpop_jkt`, `iss`
23//!    validation, and a DPoP-bound code exchange.
24//! 3. The request is retried. The response (JSON, or each event of an SSE
25//!    stream) is written to stdout. Failures become JSON-RPC errors, never
26//!    silence.
27//!
28//! ## Modules
29//!
30//! - [`proxy`]: the [`Proxy`] engine (Streamable HTTP, airlock,
31//!   DPoP, refresh).
32//! - [`auth`]: [`AuthManager`](auth::AuthManager), the OAuth flows, and
33//!   [`OidcConfig`].
34//! - [`vault`]: credential storage in the OS keychain or in memory.
35//! - [`crypto`]: DPoP keys and proofs.
36//! - [`config`]: CLI flags and environment variables.
37//! - [`logging`]: the private log directory.
38//! - [`templates`]: the default login result pages.
39//!
40//! ## Running the bridge
41//!
42//! ```no_run
43//! use mcp_airlock::{config::Config, run_with_vault, vault::Vault};
44//!
45//! # async fn demo() -> anyhow::Result<()> {
46//! let config = <Config as clap::Parser>::try_parse_from([
47//!     "mcp-airlock",
48//!     "--remote-mcp-url",
49//!     "https://mcp.example.com/mcp",
50//! ])?;
51//! // `run` picks the OS keychain; here credentials stay in memory.
52//! let vault = Vault::in_memory("example");
53//! run_with_vault(config, vault, tokio::io::stdin(), tokio::io::stdout()).await
54//! # }
55//! ```
56//!
57//! ## Sending a single request
58//!
59//! ```no_run
60//! use mcp_airlock::auth::OidcConfig;
61//! use mcp_airlock::config::AuthScheme;
62//! use mcp_airlock::proxy::Proxy;
63//! use mcp_airlock::vault::Vault;
64//! use serde_json::json;
65//!
66//! # async fn demo() -> anyhow::Result<()> {
67//! let proxy = Proxy::new(
68//!     "https://mcp.example.com/mcp",
69//!     "default_user",
70//!     OidcConfig {
71//!         client_id: "mcp-airlock".into(),
72//!         redirect_url: "http://127.0.0.1:8082/callback".into(),
73//!         ..Default::default()
74//!     },
75//!     Vault::keyring(&mcp_airlock::vault::service_name_for("https://mcp.example.com/mcp")),
76//!     "2025-11-25",
77//!     AuthScheme::Bearer,
78//! );
79//! let reply = proxy
80//!     .call(json!({
81//!         "jsonrpc": "2.0",
82//!         "id": 1,
83//!         "method": "tools/list",
84//!         "params": {"_meta": {
85//!             "io.modelcontextprotocol/protocolVersion": "2026-07-28",
86//!             "io.modelcontextprotocol/clientCapabilities": {}
87//!         }}
88//!     }))
89//!     .await?;
90//! println!("{reply:?}");
91//! # Ok(())
92//! # }
93//! ```
94#![warn(missing_docs)]
95
96pub mod auth;
97mod challenge;
98pub mod config;
99pub mod crypto;
100mod discovery;
101pub mod logging;
102mod mcp;
103mod net;
104pub mod proxy;
105pub mod templates;
106pub mod vault;
107
108use crate::auth::{OidcConfig, Timeouts};
109use crate::config::Config;
110use crate::mcp::Era;
111use crate::proxy::Proxy;
112use crate::vault::Vault;
113use std::sync::Arc;
114use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
115use tokio::sync::mpsc;
116use tokio::task::JoinSet;
117use tracing::{error, info};
118
119/// Shared result type for the crate.
120pub type Result<T> = anyhow::Result<T>;
121
122/// Runs the proxy using the vault backend selected by the environment
123/// (OS keychain unless `MCP_AIRLOCK_USE_MEMORY_VAULT` is set).
124pub async fn run<R, W>(config: Config, stdin: R, stdout: W) -> Result<()>
125where
126    R: tokio::io::AsyncRead + Unpin + Send + 'static,
127    W: tokio::io::AsyncWrite + Unpin + Send + 'static,
128{
129    let vault = Vault::from_env(&vault::service_name_for(&config.remote_mcp_url));
130    run_with_vault(config, vault, stdin, stdout).await
131}
132
133/// Checks the URL policy before anything is sent over the network.
134pub fn validate_config(config: &Config) -> Result<()> {
135    let allow = config.allow_insecure_http;
136    net::require_secure_url(&config.remote_mcp_url, "--remote-mcp-url", allow)?;
137    let optional = [
138        (&config.remote_sse_url, "--remote-sse-url"),
139        (&config.oidc_discovery_url, "--oidc-discovery-url"),
140        (&config.oidc_issuer, "--oidc-issuer"),
141        (&config.kc_auth_url, "--kc-auth-url"),
142        (&config.kc_token_url, "--kc-token-url"),
143        (&config.kc_par_url, "--kc-par-url"),
144    ];
145    for (url, what) in optional {
146        if let Some(url) = url {
147            net::require_secure_url(url, what, allow)?;
148        }
149    }
150    net::require_loopback_redirect(&config.oidc_redirect_url)?;
151    Ok(())
152}
153
154/// Runs the proxy with an explicit vault.
155pub async fn run_with_vault<R, W>(
156    config: Config,
157    vault: Vault,
158    stdin: R,
159    mut stdout: W,
160) -> Result<()>
161where
162    R: tokio::io::AsyncRead + Unpin + Send + 'static,
163    W: tokio::io::AsyncWrite + Unpin + Send + 'static,
164{
165    validate_config(&config)?;
166
167    let (stdout_tx, mut stdout_rx) = mpsc::channel::<String>(100);
168
169    // Dedicated stdout writer task
170    let stdout_handle = tokio::spawn(async move {
171        while let Some(msg) = stdout_rx.recv().await {
172            tracing::debug!(message = %msg, "Writing to stdout");
173            let mut line = msg;
174            line.push('\n');
175            if let Err(e) = stdout.write_all(line.as_bytes()).await {
176                error!("Failed to write to stdout: {:?}", e);
177                break;
178            }
179            let _ = stdout.flush().await;
180        }
181    });
182
183    let oidc_config = OidcConfig {
184        discovery_url: config.oidc_discovery_url.clone(),
185        client_id: config.oidc_client_id.clone(),
186        redirect_url: config.oidc_redirect_url.clone(),
187        auth_url_override: config.kc_auth_url.clone(),
188        token_url_override: config.kc_token_url.clone(),
189        par_url_override: config.kc_par_url.clone(),
190        template_dir: config.template_dir.clone(),
191        allow_insecure_http: config.allow_insecure_http,
192        expected_issuer: config.oidc_issuer.clone(),
193        request_offline_access: config.oidc_offline_access,
194        timeouts: Timeouts {
195            auth: std::time::Duration::from_secs(config.auth_timeout_secs),
196            ..Default::default()
197        },
198        ..Default::default()
199    };
200
201    let proxy = Proxy::new(
202        &config.remote_mcp_url,
203        &config.user_id,
204        oidc_config,
205        vault,
206        &config.mcp_protocol_version,
207        config.auth_scheme,
208    );
209    // Streamable HTTP serves the GET stream on the MCP endpoint itself.
210    let sse_url = config
211        .remote_sse_url
212        .clone()
213        .unwrap_or_else(|| config.remote_mcp_url.clone());
214
215    // Task 1: Standalone GET stream (Server -> Client). It only exists for
216    // legacy servers, so it starts once a legacy session is established.
217    // Modern servers deliver change notifications on `subscriptions/listen`.
218    let sse_proxy = proxy.clone();
219    let sse_stdout_tx = stdout_tx.clone();
220    let sse_handle = tokio::spawn(async move {
221        sse_proxy.wait_for_legacy_session().await;
222        if let Err(e) = sse_proxy.listen_sse(&sse_url, sse_stdout_tx).await {
223            error!("SSE listener failed: {:?}", e);
224        }
225    });
226
227    // Task 2: Stdio Read Loop (Client -> Server)
228    let mut reader = BufReader::new(stdin).lines();
229    let mut tasks = JoinSet::new();
230    let in_flight: InFlight = Arc::default();
231
232    info!("Ready to proxy MCP stdio messages...");
233
234    loop {
235        tokio::select! {
236            line_res = reader.next_line() => {
237                match line_res {
238                    Ok(Some(line)) => {
239                        let parsed = serde_json::from_str::<serde_json::Value>(&line).ok();
240                        if let Some(target) = parsed.as_ref().and_then(cancelled_request_key) {
241                            if let Some((handle, era)) = in_flight.lock().await.remove(&target) {
242                                info!("Client cancelled request {}; stopping it.", target);
243                                // Dropping the HTTP response closes its stream, which is
244                                // how Streamable HTTP signals cancellation. Legacy servers
245                                // also expect the notification itself.
246                                handle.abort();
247                                if era == Era::Modern {
248                                    continue;
249                                }
250                            }
251                        }
252
253                        let key = parsed.as_ref().and_then(request_key);
254                        let era = parsed.as_ref().map(mcp::era_of).unwrap_or(Era::Legacy);
255                        let proxy_task = proxy.clone();
256                        let task_stdout_tx = stdout_tx.clone();
257                        let task_in_flight = in_flight.clone();
258                        let task_key = key.clone();
259                        // Held across spawn so the task can't remove its entry
260                        // before it is inserted.
261                        let mut guard = in_flight.lock().await;
262                        let handle = tasks.spawn(async move {
263                            process_message(proxy_task, line, task_stdout_tx).await;
264                            if let Some(k) = task_key {
265                                task_in_flight.lock().await.remove(&k);
266                            }
267                        });
268                        if let Some(k) = key {
269                            guard.insert(k, (handle, era));
270                        }
271                    }
272                    Ok(None) => {
273                        info!("Stdin closed, shutting down...");
274                        break;
275                    }
276                    Err(e) => {
277                        error!("Error reading from stdin: {:?}", e);
278                        break;
279                    }
280                }
281            }
282            _ = tokio::signal::ctrl_c() => {
283                info!("Ctrl-C received, shutting down...");
284                break;
285            }
286            Some(res) = tasks.join_next(), if !tasks.is_empty() => {
287                if let Err(e) = res {
288                    if !e.is_cancelled() {
289                        error!("Proxy task failed: {:?}", e);
290                    }
291                }
292            }
293        }
294    }
295
296    // Cleanup: give in-flight requests a moment to finish, then stop. Long-lived
297    // streams (e.g. `subscriptions/listen`) would otherwise keep us alive after
298    // the client closed stdin.
299    info!("Waiting for remaining tasks to complete...");
300    sse_handle.abort();
301    let drain = async { while tasks.join_next().await.is_some() {} };
302    if tokio::time::timeout(SHUTDOWN_GRACE, drain).await.is_err() {
303        info!("Stopping requests still in flight.");
304        tasks.abort_all();
305        while tasks.join_next().await.is_some() {}
306    }
307
308    // Drop stdout_tx so the writer task can finish
309    drop(stdout_tx);
310    let _ = stdout_handle.await;
311
312    Ok(())
313}
314
315/// How long in-flight requests may run after stdin closes.
316const SHUTDOWN_GRACE: std::time::Duration = std::time::Duration::from_secs(5);
317
318/// Requests being forwarded, by JSON-RPC id, so stdio cancellations can stop them.
319type InFlight =
320    Arc<tokio::sync::Mutex<std::collections::HashMap<String, (tokio::task::AbortHandle, Era)>>>;
321
322/// The id of a request (a message with both `method` and `id`), as a map key.
323fn request_key(msg: &serde_json::Value) -> Option<String> {
324    msg.get("method")?;
325    msg.get("id").map(|id| id.to_string())
326}
327
328/// The target of a `notifications/cancelled`, as a map key.
329fn cancelled_request_key(msg: &serde_json::Value) -> Option<String> {
330    if msg.get("method")?.as_str()? != "notifications/cancelled" || msg.get("id").is_some() {
331        return None;
332    }
333    msg.pointer("/params/requestId").map(|id| id.to_string())
334}
335
336/// JSON-RPC 2.0 error codes used by the proxy.
337const PARSE_ERROR: i64 = -32700;
338const INTERNAL_ERROR: i64 = -32603;
339
340fn jsonrpc_error(id: serde_json::Value, code: i64, message: String) -> String {
341    serde_json::json!({
342        "jsonrpc": "2.0",
343        "id": id,
344        "error": { "code": code, "message": message }
345    })
346    .to_string()
347}
348
349async fn process_message(proxy: Arc<Proxy>, line: String, stdout_tx: mpsc::Sender<String>) {
350    if line.trim().is_empty() {
351        return;
352    }
353    let payload = match serde_json::from_str::<serde_json::Value>(&line) {
354        Ok(payload) => payload,
355        Err(e) => {
356            error!("Invalid JSON received on stdio: {:?}", e);
357            let msg = jsonrpc_error(serde_json::Value::Null, PARSE_ERROR, "Parse error".into());
358            let _ = stdout_tx.send(msg).await;
359            return;
360        }
361    };
362
363    // Only requests (method + id) expect a response; notifications and the
364    // client's responses to server requests must not get one.
365    let request_id = match (payload.get("method"), payload.get("id")) {
366        (Some(_), Some(id)) => Some(id.clone()),
367        _ => None,
368    };
369
370    if let Err(e) = proxy.handle_request(payload, &stdout_tx).await {
371        error!(error = ?e, "Failed to proxy request to remote server");
372        if let Some(id) = request_id {
373            let msg = jsonrpc_error(id, INTERNAL_ERROR, format!("mcp-airlock: {e:#}"));
374            let _ = stdout_tx.send(msg).await;
375        }
376    }
377}
378
379#[cfg(test)]
380mod tests {
381    use super::*;
382    use crate::config::AuthScheme;
383    use axum::response::IntoResponse;
384    use axum::{routing::post, Router};
385    use serde_json::json;
386    use tokio::io::{AsyncReadExt, AsyncWriteExt};
387
388    #[tokio::test]
389    async fn test_process_message_invalid_json() {
390        let (tx, mut rx) = mpsc::channel(1);
391        let proxy = Proxy::new(
392            "http://localhost",
393            "user",
394            OidcConfig {
395                client_id: "c".into(),
396                redirect_url: "r".into(),
397                timeouts: crate::auth::Timeouts::fast(),
398                ..Default::default()
399            },
400            Vault::in_memory("svc"),
401            "v1",
402            AuthScheme::Bearer,
403        );
404
405        process_message(proxy, "invalid json".to_string(), tx).await;
406        let resp: serde_json::Value = serde_json::from_str(&rx.try_recv().unwrap()).unwrap();
407        assert_eq!(resp["id"], serde_json::Value::Null);
408        assert_eq!(resp["error"]["code"], PARSE_ERROR);
409    }
410
411    fn proxy_for(url: &str, vault: Vault) -> Arc<Proxy> {
412        Proxy::new(
413            url,
414            "user",
415            OidcConfig {
416                client_id: "c".into(),
417                redirect_url: "r".into(),
418                timeouts: crate::auth::Timeouts::fast(),
419                ..Default::default()
420            },
421            vault,
422            "v1",
423            AuthScheme::Bearer,
424        )
425    }
426
427    async fn serve_500() -> Result<String> {
428        let app = Router::new().route(
429            "/rpc",
430            post(|| async { (axum::http::StatusCode::INTERNAL_SERVER_ERROR, "boom") }),
431        );
432        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
433        let url = format!("http://{}/rpc", listener.local_addr()?);
434        tokio::spawn(async move {
435            let _ = axum::serve(listener, app).await;
436        });
437        Ok(url)
438    }
439
440    fn vault_with_credentials() -> Result<Vault> {
441        let vault = Vault::in_memory("svc");
442        vault.store_token("user", "valid")?;
443        vault.store_dpop_key("user", &crate::crypto::DpopKey::generate().to_bytes())?;
444        Ok(vault)
445    }
446
447    #[test]
448    fn test_validate_config() {
449        use clap::Parser;
450        let parse = |extra: &[&str]| {
451            let mut args = vec!["mcp-airlock"];
452            args.extend_from_slice(extra);
453            Config::try_parse_from(args).unwrap()
454        };
455        assert!(
456            validate_config(&parse(&["--remote-mcp-url", "https://mcp.example.com/mcp"])).is_ok()
457        );
458        assert!(
459            validate_config(&parse(&["--remote-mcp-url", "http://127.0.0.1:8081/rpc"])).is_ok()
460        );
461
462        let insecure = parse(&["--remote-mcp-url", "http://mcp.example.com/mcp"]);
463        assert!(validate_config(&insecure).is_err());
464        let allowed = parse(&[
465            "--remote-mcp-url",
466            "http://mcp.example.com/mcp",
467            "--allow-insecure-http",
468        ]);
469        assert!(validate_config(&allowed).is_ok());
470
471        let bad_override = parse(&[
472            "--remote-mcp-url",
473            "https://mcp.example.com/mcp",
474            "--kc-token-url",
475            "http://as.example.com/token",
476        ]);
477        assert!(validate_config(&bad_override).is_err());
478
479        let bad_redirect = parse(&[
480            "--remote-mcp-url",
481            "https://mcp.example.com/mcp",
482            "--oidc-redirect-url",
483            "http://0.0.0.0:8082/callback",
484        ]);
485        assert!(validate_config(&bad_redirect).is_err());
486    }
487
488    #[tokio::test]
489    async fn test_process_message_failure_becomes_jsonrpc_error() -> Result<()> {
490        let (tx, mut rx) = mpsc::channel(1);
491        let proxy = proxy_for(&serve_500().await?, vault_with_credentials()?);
492
493        process_message(
494            proxy,
495            json!({"jsonrpc": "2.0", "id": "req-7", "method": "tools/list"}).to_string(),
496            tx,
497        )
498        .await;
499
500        let resp: serde_json::Value = serde_json::from_str(&rx.try_recv()?)?;
501        assert_eq!(resp["id"], "req-7");
502        assert_eq!(resp["error"]["code"], INTERNAL_ERROR);
503        let message = resp["error"]["message"].as_str().unwrap();
504        assert!(message.contains("500"), "{message}");
505        assert!(message.contains("boom"), "{message}");
506        Ok(())
507    }
508
509    #[tokio::test]
510    async fn test_process_message_client_response_gets_no_error() -> Result<()> {
511        let (tx, mut rx) = mpsc::channel(1);
512        let proxy = proxy_for(&serve_500().await?, vault_with_credentials()?);
513
514        // A response to a server-initiated request has an id but no method.
515        process_message(
516            proxy,
517            json!({"jsonrpc": "2.0", "id": 3, "result": {}}).to_string(),
518            tx,
519        )
520        .await;
521        assert!(rx.try_recv().is_err());
522        Ok(())
523    }
524
525    #[tokio::test]
526    async fn test_process_message_ignores_blank_lines() {
527        let (tx, mut rx) = mpsc::channel(1);
528        let proxy = proxy_for("http://localhost:1/rpc", Vault::in_memory("svc"));
529        process_message(proxy, "   ".to_string(), tx).await;
530        assert!(rx.try_recv().is_err());
531    }
532
533    #[tokio::test]
534    async fn test_process_message_no_id() {
535        let (tx, mut rx) = mpsc::channel(1);
536        let proxy = Proxy::new(
537            "http://localhost",
538            "user",
539            OidcConfig {
540                client_id: "c".into(),
541                redirect_url: "r".into(),
542                timeouts: crate::auth::Timeouts::fast(),
543                ..Default::default()
544            },
545            Vault::in_memory("svc"),
546            "v1",
547            AuthScheme::Bearer,
548        );
549
550        // A notification has no ID, so it shouldn't produce a response to stdout_tx
551        process_message(
552            proxy,
553            json!({"jsonrpc": "2.0", "method": "notify"}).to_string(),
554            tx,
555        )
556        .await;
557        assert!(rx.try_recv().is_err());
558    }
559
560    #[tokio::test]
561    async fn test_process_message_with_id() -> Result<()> {
562        let (tx, mut rx) = mpsc::channel(1);
563
564        let mcp_app = Router::new().route(
565            "/rpc",
566            post(|| async move { axum::Json(json!({"jsonrpc": "2.0", "id": 1, "result": "ok"})) }),
567        );
568        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
569        let addr = listener.local_addr()?;
570        let rpc_url = format!("http://127.0.0.1:{}/rpc", addr.port());
571        tokio::spawn(async move {
572            let _ = axum::serve(listener, mcp_app).await;
573        });
574
575        let vault = Vault::in_memory("test_process_message_svc");
576        let proxy = Proxy::new(
577            &rpc_url,
578            "user",
579            OidcConfig {
580                client_id: "c".into(),
581                redirect_url: "r".into(),
582                timeouts: crate::auth::Timeouts::fast(),
583                ..Default::default()
584            },
585            vault.clone(),
586            "v1",
587            AuthScheme::Bearer,
588        );
589
590        // Pre-populate vault to skip OIDC
591        vault.store_token("user", "valid")?;
592        vault.store_dpop_key("user", &crate::crypto::DpopKey::generate().to_bytes())?;
593
594        // A message with an ID should produce a response to stdout_tx
595        process_message(
596            proxy,
597            json!({"jsonrpc": "2.0", "id": 1, "method": "test"}).to_string(),
598            tx,
599        )
600        .await;
601
602        let resp = rx.recv().await.expect("Expected a response");
603        assert!(resp.contains("\"result\":\"ok\""));
604        Ok(())
605    }
606
607    /// Sets the flag when dropped, i.e. when the server stops streaming.
608    struct DropFlag(Arc<std::sync::atomic::AtomicBool>);
609    impl Drop for DropFlag {
610        fn drop(&mut self) {
611            self.0.store(true, std::sync::atomic::Ordering::SeqCst);
612        }
613    }
614
615    /// A server whose `tools/call` streams forever (until the client
616    /// disconnects) and which records every message POSTed to it.
617    async fn endless_stream_server() -> Result<(
618        String,
619        Arc<std::sync::atomic::AtomicBool>,
620        Arc<std::sync::Mutex<Vec<serde_json::Value>>>,
621    )> {
622        use axum::response::sse::{Event, Sse};
623        let closed = Arc::new(std::sync::atomic::AtomicBool::new(false));
624        let posted: Arc<std::sync::Mutex<Vec<serde_json::Value>>> = Arc::default();
625        let (c, p) = (closed.clone(), posted.clone());
626        let app = Router::new().route(
627            "/rpc",
628            post(move |axum::Json(body): axum::Json<serde_json::Value>| {
629                let (c, p) = (c.clone(), p.clone());
630                async move {
631                    p.lock().unwrap().push(body.clone());
632                    if body["method"] != "tools/call" {
633                        return axum::http::StatusCode::ACCEPTED.into_response();
634                    }
635                    let guard = DropFlag(c);
636                    let stream = futures::stream::unfold(guard, |g| async move {
637                        tokio::time::sleep(std::time::Duration::from_millis(20)).await;
638                        let progress = json!({"jsonrpc": "2.0", "method": "notifications/progress",
639                            "params": {"progressToken": "t", "progress": 1}});
640                        Some((
641                            Ok::<_, std::convert::Infallible>(
642                                Event::default().data(progress.to_string()),
643                            ),
644                            g,
645                        ))
646                    });
647                    Sse::new(stream).into_response()
648                }
649            }),
650        );
651        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
652        let url = format!("http://{}/rpc", listener.local_addr()?);
653        tokio::spawn(async move {
654            let _ = axum::serve(listener, app).await;
655        });
656        Ok((url, closed, posted))
657    }
658
659    async fn run_cancellation(
660        call: serde_json::Value,
661    ) -> Result<(bool, Vec<serde_json::Value>, String)> {
662        let (url, closed, posted) = endless_stream_server().await?;
663        let config = <Config as clap::Parser>::try_parse_from([
664            "mcp-airlock",
665            "--remote-mcp-url",
666            &url,
667            "--oidc-redirect-url",
668            "http://127.0.0.1:1/callback",
669        ])?;
670        let (mut client_out, server_out) = tokio::io::duplex(64 * 1024);
671        let (mut client_in, server_in) = tokio::io::duplex(64 * 1024);
672        let run = tokio::spawn(run_with_vault(
673            config,
674            vault_with_credentials()?,
675            server_in,
676            server_out,
677        ));
678
679        client_in.write_all(format!("{call}\n").as_bytes()).await?;
680        // Cancel only once the server is streaming: a cancellation that beats the
681        // request there leaves nothing to close. A fixed delay is not enough when
682        // the tests run slowly (e.g. under coverage instrumentation).
683        let deadline = std::time::Instant::now() + std::time::Duration::from_secs(10);
684        while posted.lock().unwrap().is_empty() && std::time::Instant::now() < deadline {
685            tokio::time::sleep(std::time::Duration::from_millis(10)).await;
686        }
687        let cancel = json!({"jsonrpc": "2.0", "method": "notifications/cancelled",
688            "params": {"requestId": call["id"], "reason": "user"}});
689        client_in
690            .write_all(format!("{cancel}\n").as_bytes())
691            .await?;
692
693        let deadline = std::time::Instant::now() + std::time::Duration::from_secs(3);
694        while !closed.load(std::sync::atomic::Ordering::SeqCst)
695            && std::time::Instant::now() < deadline
696        {
697            tokio::time::sleep(std::time::Duration::from_millis(10)).await;
698        }
699        let was_closed = closed.load(std::sync::atomic::Ordering::SeqCst);
700        drop(client_in);
701        tokio::time::timeout(std::time::Duration::from_secs(10), run).await???;
702
703        let mut out = String::new();
704        client_out.read_to_string(&mut out).await?;
705        let posted = posted.lock().unwrap().clone();
706        Ok((was_closed, posted, out))
707    }
708
709    #[tokio::test]
710    async fn test_modern_cancellation_closes_the_stream() -> Result<()> {
711        let call = json!({"jsonrpc": "2.0", "id": 5, "method": "tools/call",
712            "params": {"name": "slow", "arguments": {},
713                "_meta": {"io.modelcontextprotocol/protocolVersion": "2026-07-28",
714                          "io.modelcontextprotocol/clientCapabilities": {}}}});
715        let (closed, posted, out) = run_cancellation(call).await?;
716        assert!(closed, "the server must see the response stream close");
717        // Streamable HTTP defines no cancellation notification: none is POSTed.
718        assert_eq!(posted.len(), 1);
719        // Progress may have been forwarded, but never a response for id 5.
720        for line in out.lines() {
721            let msg: serde_json::Value = serde_json::from_str(line)?;
722            assert!(msg.get("id").is_none(), "unexpected response: {line}");
723        }
724        Ok(())
725    }
726
727    #[tokio::test]
728    async fn test_legacy_cancellation_also_forwards_the_notification() -> Result<()> {
729        let call = json!({"jsonrpc": "2.0", "id": "c-1", "method": "tools/call",
730            "params": {"name": "slow", "arguments": {}}});
731        let (closed, posted, _out) = run_cancellation(call).await?;
732        assert!(closed);
733        assert_eq!(posted.len(), 2);
734        assert_eq!(posted[1]["method"], "notifications/cancelled");
735        assert_eq!(posted[1]["params"]["requestId"], "c-1");
736        Ok(())
737    }
738
739    #[test]
740    fn test_request_and_cancel_keys() {
741        assert_eq!(
742            request_key(&json!({"id": 1, "method": "x"})),
743            Some("1".into())
744        );
745        assert_eq!(
746            request_key(&json!({"id": "a", "method": "x"})),
747            Some("\"a\"".into())
748        );
749        assert_eq!(request_key(&json!({"id": 1, "result": {}})), None);
750        assert_eq!(request_key(&json!({"method": "notify"})), None);
751        let cancel = json!({"method": "notifications/cancelled", "params": {"requestId": "a"}});
752        assert_eq!(cancelled_request_key(&cancel), Some("\"a\"".into()));
753        assert_eq!(cancelled_request_key(&json!({"method": "other"})), None);
754    }
755
756    #[tokio::test]
757    async fn test_run_minimal() -> Result<()> {
758        let (mut client_out_rx, server_out_tx) = tokio::io::duplex(1024);
759        let (mut client_in_tx, server_in_rx) = tokio::io::duplex(1024);
760
761        let mcp_app = Router::new().route(
762            "/rpc",
763            post(|| async move { axum::Json(json!({"jsonrpc": "2.0", "id": 1, "result": "ok"})) }),
764        );
765        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
766        let addr = listener.local_addr()?;
767        let rpc_url = format!("http://127.0.0.1:{}/rpc", addr.port());
768        tokio::spawn(async move {
769            let _ = axum::serve(listener, mcp_app).await;
770        });
771
772        let config = Config {
773            remote_mcp_url: rpc_url,
774            remote_sse_url: Some(format!("http://127.0.0.1:{}/sse", addr.port())),
775            user_id: "test-user".into(),
776            oidc_discovery_url: None,
777            oidc_client_id: "client".into(),
778            oidc_redirect_url: "http://localhost:1/callback".into(),
779            kc_auth_url: Some("http://localhost:1/auth".into()),
780            kc_token_url: Some("http://localhost:1/token".into()),
781            kc_par_url: Some("http://localhost:1/par".into()),
782            log_level: "info".into(),
783            log_dir: None,
784            allow_insecure_http: false,
785            oidc_issuer: None,
786            oidc_offline_access: false,
787            template_dir: None,
788            mcp_protocol_version: "2025-11-25".into(),
789            auth_scheme: AuthScheme::Bearer,
790            auth_timeout_secs: 300,
791        };
792
793        // Pre-populate vault to skip OIDC
794        let vault = Vault::in_memory("mcp-airlock");
795        vault.store_token("test-user", "valid")?;
796        vault.store_dpop_key("test-user", &crate::crypto::DpopKey::generate().to_bytes())?;
797
798        let run_handle = tokio::spawn(async move {
799            run_with_vault(config, vault, server_in_rx, server_out_tx).await
800        });
801
802        // Send a message
803        client_in_tx
804            .write_all(b"{\"jsonrpc\": \"2.0\", \"id\": 1, \"method\": \"test\"}\n")
805            .await?;
806
807        // Wait for response
808        let mut buf = [0u8; 1024];
809        let n = client_out_rx.read(&mut buf).await?;
810        let resp = String::from_utf8_lossy(&buf[..n]);
811        assert!(resp.contains("\"result\":\"ok\""));
812
813        // Send another message
814        client_in_tx
815            .write_all(b"{\"jsonrpc\": \"2.0\", \"id\": 2, \"method\": \"test2\"}\n")
816            .await?;
817        let n = client_out_rx.read(&mut buf).await?;
818        let resp = String::from_utf8_lossy(&buf[..n]);
819        assert!(resp.contains("\"result\":\"ok\""));
820
821        // Close stdin to trigger shutdown
822        drop(client_in_tx);
823
824        let res = tokio::time::timeout(std::time::Duration::from_secs(2), run_handle).await??;
825        assert!(res.is_ok());
826        Ok(())
827    }
828}