Skip to main content

mcp_airlock/
proxy.rs

1//! # Transparent Layer 7 Bridge (Airlock)
2//!
3//! This module implements the core proxying logic between the AI client's stdio and the
4//! remote MCP server's HTTP/SSE interface. It features the "Airlock" mechanism
5//! for transparently handling authentication challenges without interrupting the client connection.
6
7use crate::auth::{AuthManager, OidcConfig};
8use crate::challenge::WwwAuthenticate;
9use crate::config::AuthScheme;
10use crate::crypto::DpopKey;
11use crate::mcp::{self, Era};
12use crate::vault::Vault;
13use crate::Result;
14use anyhow::Context;
15use rand::Rng;
16use reqwest::header::{ACCEPT, AUTHORIZATION, CONTENT_TYPE};
17use reqwest::{Client, StatusCode};
18use serde_json::Value;
19use std::sync::Arc;
20use tokio::sync::{mpsc, watch, Mutex, RwLock};
21use tracing::{error, info, warn};
22use url::Url;
23
24/// The main proxy engine that manages the connection and authentication state.
25pub struct Proxy {
26    /// The HTTP client used for proxying requests.
27    http_client: Client,
28    /// The base URL of the remote MCP server.
29    remote_url: String,
30    /// Receiver for the suspension state (Airlock status).
31    suspension_rx: watch::Receiver<bool>,
32    /// Sender for the suspension state (Airlock status).
33    suspension_tx: watch::Sender<bool>,
34    /// Secure vault for storing tokens and keys.
35    vault: Vault,
36    /// Unique identifier for the current user.
37    user_id: String,
38    /// OIDC configuration and metadata.
39    oidc_config: OidcConfig,
40    /// Shared authentication manager (lazy-loaded).
41    pub auth_manager: Arc<RwLock<Option<Arc<AuthManager>>>>,
42    /// The MCP protocol version to use.
43    protocol_version: String,
44    /// The authentication scheme (Bearer or DPoP).
45    auth_scheme: AuthScheme,
46    /// Current session ID (for persistent SSE).
47    session_id: Mutex<Option<String>>,
48    /// Serializes re-authentication attempts.
49    reauth_mutex: Mutex<()>,
50    /// Outcome of past re-authentication attempts.
51    reauth_state: Mutex<ReauthState>,
52    /// Credential generation: bumped on every successful re-authentication.
53    reauth_count: Arc<std::sync::atomic::AtomicU64>,
54    /// Set once a legacy session exists (an `initialize` result or a session
55    /// id); only legacy servers offer the standalone GET stream.
56    legacy_session_tx: watch::Sender<bool>,
57    /// Latest DPoP nonce provided by the remote server (RFC 9449 §9).
58    rs_nonce: std::sync::Mutex<Option<String>>,
59    /// Credential generation for which a proactive refresh was already tried.
60    proactive_refresh_gen: std::sync::Mutex<Option<u64>>,
61    /// Version negotiated by a legacy `initialize` handshake.
62    negotiated_version: std::sync::Mutex<Option<String>>,
63    /// `x-mcp-header` annotations of the tools seen in `tools/list` results.
64    tool_headers: std::sync::Mutex<std::collections::HashMap<String, mcp::ToolHeaders>>,
65}
66
67/// A fresh token rejected within this window means re-authenticating again
68/// would loop.
69const AUTH_LOOP_WINDOW: std::time::Duration = std::time::Duration::from_secs(5);
70
71/// Why a request needs new credentials.
72#[derive(Clone, Copy, Debug, PartialEq, Eq)]
73pub enum ReauthReason {
74    /// 401: the token is missing, expired or revoked. A refresh may fix it.
75    Unauthorized,
76    /// 403 `insufficient_scope`: more scopes are needed, which takes a new login.
77    StepUp,
78    /// The token is about to expire: refresh it if possible, never log in.
79    Expiring,
80}
81
82/// What the proxy remembers about re-authentication attempts.
83#[derive(Default)]
84struct ReauthState {
85    /// Number of attempts made, successful or not.
86    attempts: u64,
87    /// Why the last attempt failed, if it did.
88    last_error: Option<String>,
89    /// The last successful attempt.
90    last_success: Option<LastSuccess>,
91}
92
93struct LastSuccess {
94    at: std::time::Instant,
95    /// Scopes requested by that attempt.
96    scopes: Option<Vec<String>>,
97    /// Whether the token came from a refresh rather than a browser login.
98    via_refresh: bool,
99}
100
101/// How one round trip of [`Proxy::send_once`] ended.
102enum Outcome {
103    /// The reply (if any) was forwarded.
104    Done,
105    /// The server rejected the request's headers (`-32020`); not forwarded yet.
106    HeaderMismatch(Value),
107}
108
109/// Refresh this long before the access token expires.
110const EXPIRY_MARGIN_SECS: u64 = 30;
111
112/// A token and the DPoP key it is bound to.
113struct Credentials {
114    token: String,
115    key: DpopKey,
116    /// When the token expires (Unix seconds), if known.
117    expires_at: Option<u64>,
118}
119
120impl Credentials {
121    fn expires_soon(&self) -> bool {
122        let now = std::time::SystemTime::now()
123            .duration_since(std::time::UNIX_EPOCH)
124            .map(|d| d.as_secs())
125            .unwrap_or_default();
126        self.expires_at
127            .is_some_and(|at| at <= now + EXPIRY_MARGIN_SECS)
128    }
129}
130
131fn validate_resource_metadata(metadata_url: Option<&str>, remote_url: &str) -> Option<String> {
132    let url_str = metadata_url?;
133
134    let m_url = match Url::parse(url_str) {
135        Ok(u) => u,
136        Err(e) => {
137            warn!("Failed to parse resource_metadata URL ({}): {}", url_str, e);
138            return None;
139        }
140    };
141
142    let r_url = match Url::parse(remote_url) {
143        Ok(u) => u,
144        Err(e) => {
145            warn!("Failed to parse remote URL ({}): {}", remote_url, e);
146            return None;
147        }
148    };
149
150    // Same origin (scheme, host and port), so a challenge can neither point us
151    // at another host (SSRF) nor downgrade https to http.
152    if m_url.origin() != r_url.origin() {
153        warn!(
154            "SSRF Prevention: resource_metadata origin {} does not match remote origin {}. Rejecting.",
155            m_url.origin().ascii_serialization(),
156            r_url.origin().ascii_serialization()
157        );
158        return None;
159    }
160
161    Some(url_str.to_string())
162}
163
164impl Proxy {
165    /// Creates a new Proxy instance.
166    pub fn new(
167        remote_url: &str,
168        user_id: &str,
169        oidc_config: OidcConfig,
170        vault: Vault,
171        protocol_version: &str,
172        auth_scheme: AuthScheme,
173    ) -> Arc<Self> {
174        let (tx, rx) = watch::channel(false);
175        let http_client = Client::builder()
176            .connect_timeout(crate::auth::CONNECT_TIMEOUT)
177            .build()
178            .expect("Failed to build HTTP client");
179        Arc::new(Self {
180            http_client,
181            remote_url: remote_url.to_string(),
182            suspension_rx: rx,
183            suspension_tx: tx,
184            vault,
185            user_id: user_id.to_string(),
186            oidc_config,
187            auth_manager: Arc::new(RwLock::new(None)),
188            protocol_version: protocol_version.to_string(),
189            auth_scheme,
190            session_id: Mutex::new(None),
191            reauth_mutex: Mutex::new(()),
192            reauth_state: Mutex::new(ReauthState::default()),
193            reauth_count: Arc::new(std::sync::atomic::AtomicU64::new(0)),
194            legacy_session_tx: watch::channel(false).0,
195            rs_nonce: std::sync::Mutex::new(None),
196            proactive_refresh_gen: std::sync::Mutex::new(None),
197            negotiated_version: std::sync::Mutex::new(None),
198            tool_headers: std::sync::Mutex::default(),
199        })
200    }
201
202    /// Loads the token and its DPoP key. A token without its key is unusable
203    /// and treated as no credentials.
204    fn load_credentials(&self) -> Result<Option<Credentials>> {
205        let Some(token) = self.vault.get_token(&self.user_id)? else {
206            return Ok(None);
207        };
208        match self.vault.get_dpop_key(&self.user_id)? {
209            Some(bytes) => Ok(Some(Credentials {
210                token,
211                key: DpopKey::from_bytes(&bytes)?,
212                expires_at: self
213                    .vault
214                    .get_meta(&self.user_id)?
215                    .and_then(|m| m.expires_at),
216            })),
217            None => {
218                warn!("Stored token has no DPoP key; it will not be used.");
219                Ok(None)
220            }
221        }
222    }
223
224    /// Remembers the server's `DPoP-Nonce`, if it sent one. Returns whether it did.
225    fn update_rs_nonce(&self, headers: &reqwest::header::HeaderMap) -> bool {
226        match crate::crypto::dpop_nonce(headers) {
227            Some(n) => {
228                if let Ok(mut slot) = self.rs_nonce.lock() {
229                    *slot = Some(n);
230                }
231                true
232            }
233            None => false,
234        }
235    }
236
237    /// Whether a response asks for a (new) DPoP nonce (RFC 9449 §9).
238    fn wants_dpop_nonce(status: StatusCode, headers: &reqwest::header::HeaderMap) -> bool {
239        status == StatusCode::UNAUTHORIZED
240            && WwwAuthenticate::parse(headers).has_error("use_dpop_nonce")
241    }
242
243    /// Returns true the first time it is called for `gen`.
244    fn claim_proactive_refresh(&self, gen: u64) -> bool {
245        match self.proactive_refresh_gen.lock() {
246            Ok(mut slot) if *slot != Some(gen) => {
247                *slot = Some(gen);
248                true
249            }
250            _ => false,
251        }
252    }
253
254    /// The current credential generation.
255    fn generation(&self) -> u64 {
256        self.reauth_count.load(std::sync::atomic::Ordering::SeqCst)
257    }
258
259    async fn ensure_auth_manager(&self, metadata_url: Option<&str>) -> Result<Arc<AuthManager>> {
260        {
261            let lock = self.auth_manager.read().await;
262            if let Some(am) = lock.as_ref() {
263                return Ok(am.clone());
264            }
265        }
266
267        let mut lock = self.auth_manager.write().await;
268        // Re-check after acquiring write lock
269        if let Some(am) = lock.as_ref() {
270            return Ok(am.clone());
271        }
272
273        info!(
274            "Initializing AuthManager (resource_metadata: {:?})...",
275            metadata_url
276        );
277
278        let am = AuthManager::discover(
279            self.oidc_config.clone(),
280            self.remote_url.clone(), // This is the 'resource'
281            self.vault.clone(),
282            metadata_url,
283        )
284        .await?;
285
286        let am_shared = Arc::new(am);
287        *lock = Some(am_shared.clone());
288        Ok(am_shared)
289    }
290
291    /// The `MCP-Protocol-Version` for a message: the `_meta` version of a
292    /// modern request, the requested version of `initialize`, the version a
293    /// legacy handshake negotiated, or the configured fallback.
294    fn protocol_version_for(&self, payload: Option<&Value>) -> String {
295        if let Some(payload) = payload {
296            if let Some(v) = mcp::meta_protocol_version(payload) {
297                return v.to_string();
298            }
299            if mcp::method_of(payload) == Some("initialize") {
300                if let Some(v) = payload
301                    .pointer("/params/protocolVersion")
302                    .and_then(Value::as_str)
303                {
304                    return v.to_string();
305                }
306            }
307        }
308        self.negotiated_version
309            .lock()
310            .ok()
311            .and_then(|v| v.clone())
312            .unwrap_or_else(|| self.protocol_version.clone())
313    }
314
315    /// Inspects a message the server sent in reply to `request` before it is
316    /// forwarded, and returns the message to forward.
317    fn observe_response(&self, request: &Value, mut message: Value) -> Value {
318        if mcp::method_of(request) == Some("initialize") && message.get("id") == request.get("id") {
319            if let Some(v) = message
320                .pointer("/result/protocolVersion")
321                .and_then(Value::as_str)
322            {
323                info!("Legacy MCP session negotiated protocol version {}", v);
324                if let Ok(mut slot) = self.negotiated_version.lock() {
325                    *slot = Some(v.to_string());
326                }
327                self.legacy_session_tx.send_replace(true);
328            }
329        }
330        if mcp::method_of(request) == Some("tools/list") && message.get("id") == request.get("id") {
331            if let Some(tools) = message
332                .pointer_mut("/result/tools")
333                .and_then(Value::as_array_mut)
334            {
335                self.record_tools(tools);
336            }
337        }
338        message
339    }
340
341    /// Caches the `x-mcp-header` annotations of `tools` and removes the tools
342    /// whose annotations are invalid, as Streamable HTTP clients must.
343    fn record_tools(&self, tools: &mut Vec<Value>) {
344        let Ok(mut cache) = self.tool_headers.lock() else {
345            return;
346        };
347        tools.retain(|tool| {
348            let name = tool.get("name").and_then(Value::as_str).unwrap_or_default();
349            let schema = tool.get("inputSchema").unwrap_or(&Value::Null);
350            match mcp::tool_header_annotations(schema) {
351                Ok(headers) => {
352                    if headers.is_empty() {
353                        cache.remove(name);
354                    } else {
355                        cache.insert(name.to_string(), headers);
356                    }
357                    true
358                }
359                Err(reason) => {
360                    warn!("Rejecting tool '{}' from tools/list: {}", name, reason);
361                    cache.remove(name);
362                    false
363                }
364            }
365        });
366    }
367
368    /// `Mcp-Param-*` headers for a `tools/call`, from the cached annotations.
369    fn tool_param_headers(&self, payload: &Value) -> Vec<(String, String)> {
370        if mcp::method_of(payload) != Some("tools/call") {
371            return Vec::new();
372        }
373        let Some(name) = payload.pointer("/params/name").and_then(Value::as_str) else {
374            return Vec::new();
375        };
376        let cache = match self.tool_headers.lock() {
377            Ok(cache) => cache,
378            Err(_) => return Vec::new(),
379        };
380        match cache.get(name) {
381            Some(headers) => mcp::param_headers(payload.pointer("/params/arguments"), headers),
382            None => Vec::new(),
383        }
384    }
385
386    /// Re-reads `tools/list` (all pages, up to a limit) to refresh the cached
387    /// annotations, using the protocol metadata of `request`.
388    async fn refresh_tool_headers(&self, request: &Value) -> Result<()> {
389        const MAX_PAGES: usize = 10;
390        let mut cursor: Option<Value> = None;
391        for page in 0..MAX_PAGES {
392            let mut params = serde_json::Map::new();
393            if let Some(meta) = request.pointer("/params/_meta") {
394                params.insert("_meta".into(), meta.clone());
395            }
396            if let Some(c) = cursor.take() {
397                params.insert("cursor".into(), c);
398            }
399            let list = serde_json::json!({
400                "jsonrpc": "2.0",
401                "id": format!("mcp-airlock-tools-list-{page}"),
402                "method": "tools/list",
403                "params": params,
404            });
405            let (tx, mut rx) = mpsc::channel(16);
406            let outcome = self.send_once(&list, &tx).await?;
407            drop(tx);
408            if let Outcome::HeaderMismatch(v) = outcome {
409                anyhow::bail!("tools/list was rejected: {}", v);
410            }
411            let mut next = None;
412            while let Some(m) = rx.recv().await {
413                if let Ok(v) = serde_json::from_str::<Value>(&m) {
414                    if v.get("id") == list.get("id") {
415                        next = v.pointer("/result/nextCursor").cloned();
416                    }
417                }
418            }
419            match next {
420                Some(c) if !c.is_null() => cursor = Some(c),
421                _ => return Ok(()),
422            }
423        }
424        Ok(())
425    }
426
427    /// Builds a request to the remote server with the MCP headers and, when a
428    /// token is available, the Authorization and DPoP headers. `payload` is the
429    /// JSON-RPC message of a POST.
430    async fn build_request(
431        &self,
432        method: reqwest::Method,
433        url: &str,
434        credentials: Option<&Credentials>,
435        payload: Option<&Value>,
436    ) -> Result<reqwest::RequestBuilder> {
437        let mut request = self
438            .http_client
439            .request(method.clone(), url)
440            .header("MCP-Protocol-Version", self.protocol_version_for(payload));
441        if let Some(payload) = payload {
442            for (name, value) in mcp::standard_headers(payload) {
443                request = request.header(name, value);
444            }
445            for (name, value) in self.tool_param_headers(payload) {
446                request = request.header(name, value);
447            }
448        }
449
450        if let Some(Credentials { token, key, .. }) = credentials {
451            let nonce = self.rs_nonce.lock().ok().and_then(|n| n.clone());
452            let dpop_proof =
453                key.generate_proof_with_ath(method.as_str(), url, Some(token), nonce.as_deref())?;
454            let auth_header = match self.auth_scheme {
455                AuthScheme::Bearer => format!("Bearer {}", token),
456                AuthScheme::Dpop => format!("DPoP {}", token),
457            };
458            request = request
459                .header(AUTHORIZATION, auth_header)
460                .header("DPoP", dpop_proof);
461        }
462
463        // Sessions only exist in the legacy era.
464        let legacy = payload.is_none_or(|p| mcp::era_of(p) == Era::Legacy);
465        if legacy {
466            if let Some(s) = &*self.session_id.lock().await {
467                request = request.header("MCP-Session-Id", s);
468            }
469        }
470        Ok(request)
471    }
472
473    async fn execute_request(
474        &self,
475        credentials: Option<&Credentials>,
476        payload: &Value,
477    ) -> Result<reqwest::Response> {
478        if credentials.is_none() {
479            info!(
480                "No credentials for user, sending unauthenticated request to trigger discovery..."
481            );
482        }
483        let request = self
484            .build_request(
485                reqwest::Method::POST,
486                &self.remote_url,
487                credentials,
488                Some(payload),
489            )
490            .await?
491            .header(ACCEPT, "application/json, text/event-stream");
492        Ok(request.json(payload).send().await?)
493    }
494
495    /// Re-authenticates according to a `WWW-Authenticate` challenge.
496    async fn reauth_for_challenge(
497        &self,
498        challenge: &WwwAuthenticate,
499        observed_gen: u64,
500        reason: ReauthReason,
501    ) -> Result<()> {
502        // Without a (valid) resource_metadata, discovery falls back to the well-known URIs.
503        let metadata_url =
504            validate_resource_metadata(challenge.resource_metadata(), &self.remote_url);
505        self.trigger_reauth(
506            observed_gen,
507            metadata_url.as_deref(),
508            challenge.scope(),
509            reason,
510        )
511        .await
512    }
513
514    /// Handles 401 (expired/missing token) and 403 `insufficient_scope` (step-up).
515    /// Returns `true` when the request should be retried.
516    async fn handle_auth_challenge(
517        &self,
518        response: &reqwest::Response,
519        observed_gen: u64,
520    ) -> Result<bool> {
521        let status = response.status();
522        if status != StatusCode::UNAUTHORIZED && status != StatusCode::FORBIDDEN {
523            return Ok(false);
524        }
525        let challenge = WwwAuthenticate::parse(response.headers());
526        let reason = if status == StatusCode::UNAUTHORIZED {
527            warn!("401 Unauthorized received. Activating Airlock suspension...");
528            ReauthReason::Unauthorized
529        } else if challenge.error() == Some("insufficient_scope") {
530            warn!(
531                "403 Forbidden (insufficient_scope) received. Triggering step-up authentication..."
532            );
533            ReauthReason::StepUp
534        } else {
535            return Ok(false);
536        };
537        self.reauth_for_challenge(&challenge, observed_gen, reason)
538            .await?;
539        Ok(true)
540    }
541
542    /// Primary entry point for the stdio -> HTTP bridge.
543    ///
544    /// Sends one JSON-RPC message to the remote server (attaching DPoP-bound
545    /// tokens and managing the Airlock) and writes every message the server
546    /// returns for it to `out`: a JSON body, or each event of a
547    /// `text/event-stream` response (MCP Streamable HTTP).
548    pub async fn handle_request(&self, payload: Value, out: &mpsc::Sender<String>) -> Result<()> {
549        match self.send_once(&payload, out).await? {
550            Outcome::Done => Ok(()),
551            Outcome::HeaderMismatch(error) => {
552                // The tool's x-mcp-header annotations probably changed: re-read
553                // tools/list and retry once with the new headers.
554                warn!("Server reported a header mismatch; refreshing tools/list and retrying.");
555                if let Err(e) = self.refresh_tool_headers(&payload).await {
556                    warn!("Could not refresh tools/list: {:#}", e);
557                }
558                match self.send_once(&payload, out).await? {
559                    Outcome::Done => Ok(()),
560                    Outcome::HeaderMismatch(_) => {
561                        let _ = out.send(error.to_string()).await;
562                        Ok(())
563                    }
564                }
565            }
566        }
567    }
568
569    /// Sends `payload` once (with auth and nonce retries) and forwards the
570    /// reply, except a `HeaderMismatch` error which is returned to the caller.
571    async fn send_once(&self, payload: &Value, out: &mpsc::Sender<String>) -> Result<Outcome> {
572        let max_retries = 2;
573        let mut retry_count = 0;
574        let mut loaded_gen: Option<u64> = None;
575        let mut credentials = None;
576        let mut nonce_retried = false;
577
578        let response = loop {
579            if retry_count > max_retries {
580                error!("Maximum retry attempts reached for request. Aborting to prevent infinite loop.");
581                anyhow::bail!("Maximum retry attempts reached");
582            }
583
584            self.wait_for_airlock().await?;
585            let mut gen = self.generation();
586            if loaded_gen != Some(gen) {
587                credentials = self.load_credentials()?;
588                loaded_gen = Some(gen);
589            }
590
591            // Renew a token that is about to expire instead of waiting for a 401
592            // (once per credential generation, so a failing refresh isn't retried
593            // on every request).
594            if credentials.as_ref().is_some_and(Credentials::expires_soon)
595                && self.claim_proactive_refresh(gen)
596            {
597                info!("Access token expires soon; refreshing it.");
598                if let Err(e) = self
599                    .trigger_reauth(gen, None, None, ReauthReason::Expiring)
600                    .await
601                {
602                    warn!("Proactive refresh failed: {:#}", e);
603                }
604                gen = self.generation();
605                if loaded_gen != Some(gen) {
606                    credentials = self.load_credentials()?;
607                    loaded_gen = Some(gen);
608                }
609            }
610
611            let response = self.execute_request(credentials.as_ref(), payload).await?;
612            let got_nonce = self.update_rs_nonce(response.headers());
613            if Self::wants_dpop_nonce(response.status(), response.headers()) {
614                // Not an auth failure: resend once with the nonce, outside the airlock.
615                if got_nonce && !nonce_retried {
616                    info!("Remote server requires a DPoP nonce; retrying.");
617                    nonce_retried = true;
618                    continue;
619                }
620                anyhow::bail!("Remote server rejected the DPoP nonce (use_dpop_nonce)");
621            }
622            if self.handle_auth_challenge(&response, gen).await? {
623                retry_count += 1;
624                continue;
625            }
626            break response;
627        };
628
629        let status = response.status();
630
631        // A legacy request carrying a session id that gets a bare 404 means the
632        // session is gone; the client has to start a new one with `initialize`.
633        // (A modern 404 carries a JSON-RPC error, e.g. method not found.)
634        if status == StatusCode::NOT_FOUND && mcp::era_of(payload) == Era::Legacy {
635            let body = response.bytes().await?;
636            let jsonrpc_error = serde_json::from_slice::<Value>(&body)
637                .ok()
638                .filter(|v| v.get("jsonrpc").is_some() && v.get("error").is_some());
639            if let Some(value) = jsonrpc_error {
640                let _ = out.send(value.to_string()).await;
641                return Ok(Outcome::Done);
642            }
643            if let Some(old) = self.session_id.lock().await.take() {
644                warn!("MCP session {} expired (HTTP 404).", old);
645                anyhow::bail!("MCP session expired; re-initialize the connection");
646            }
647            anyhow::bail!("Remote MCP server returned HTTP 404");
648        }
649
650        if let Some(sid) = response
651            .headers()
652            .get("mcp-session-id")
653            .and_then(|h| h.to_str().ok())
654        {
655            let mut sid_lock = self.session_id.lock().await;
656            if sid_lock.as_deref() != Some(sid) {
657                info!("New MCP Session ID captured: {}", sid);
658                *sid_lock = Some(sid.to_string());
659            }
660            self.legacy_session_tx.send_replace(true);
661        }
662
663        if status == StatusCode::ACCEPTED || status == StatusCode::NO_CONTENT {
664            return Ok(Outcome::Done);
665        }
666
667        let is_event_stream = response
668            .headers()
669            .get(CONTENT_TYPE)
670            .and_then(|v| v.to_str().ok())
671            .is_some_and(|v| v.starts_with("text/event-stream"));
672
673        if status.is_success() && is_event_stream {
674            self.forward_event_stream(response, payload, out).await?;
675            return Ok(Outcome::Done);
676        }
677
678        let body = response.bytes().await?;
679        if status.is_success() {
680            if body.iter().all(u8::is_ascii_whitespace) {
681                return Ok(Outcome::Done);
682            }
683            let value: Value = serde_json::from_slice(&body)
684                .context("Remote MCP server returned a body that is not JSON")?;
685            let value = self.observe_response(payload, value);
686            let _ = out.send(value.to_string()).await;
687            return Ok(Outcome::Done);
688        }
689
690        // Some servers put a JSON-RPC error in a non-2xx response: pass it on as is.
691        if let Ok(value) = serde_json::from_slice::<Value>(&body) {
692            if value.get("jsonrpc").is_some() && value.get("error").is_some() {
693                let mismatch = status == StatusCode::BAD_REQUEST
694                    && mcp::method_of(payload) == Some("tools/call")
695                    && value.pointer("/error/code").and_then(Value::as_i64)
696                        == Some(mcp::HEADER_MISMATCH);
697                if mismatch {
698                    return Ok(Outcome::HeaderMismatch(value));
699                }
700                let _ = out.send(value.to_string()).await;
701                return Ok(Outcome::Done);
702            }
703        }
704        let text = String::from_utf8_lossy(&body);
705        let snippet: String = text.chars().take(200).collect();
706        anyhow::bail!(
707            "Remote MCP server returned HTTP {}: {}",
708            status,
709            snippet.trim()
710        )
711    }
712
713    /// Sends one request and returns the server's response to it.
714    ///
715    /// Other messages the server streams back (e.g. progress notifications) are
716    /// dropped. Returns `Ok(None)` when nothing comes back, as for notifications.
717    pub async fn call(&self, payload: Value) -> Result<Option<Value>> {
718        let id = payload.get("id").cloned();
719        let (tx, mut rx) = mpsc::channel::<String>(16);
720        let (res, messages) = tokio::join!(
721            async move { self.handle_request(payload, &tx).await },
722            async {
723                let mut messages = Vec::new();
724                while let Some(m) = rx.recv().await {
725                    messages.push(m);
726                }
727                messages
728            }
729        );
730        res?;
731        let mut parsed: Vec<Value> = messages
732            .iter()
733            .filter_map(|m| serde_json::from_str(m).ok())
734            .collect();
735        let pos = parsed
736            .iter()
737            .position(|m| id.is_some() && m.get("id") == id.as_ref());
738        Ok(match pos {
739            Some(i) => Some(parsed.swap_remove(i)),
740            None => parsed.pop(),
741        })
742    }
743
744    /// Forwards each event of a `text/event-stream` POST response to `out`.
745    ///
746    /// Fails if the stream ends before the response to the request arrived, so
747    /// that the client gets an error instead of waiting forever.
748    async fn forward_event_stream(
749        &self,
750        response: reqwest::Response,
751        request: &Value,
752        out: &mpsc::Sender<String>,
753    ) -> Result<()> {
754        use eventsource_stream::Eventsource;
755        use futures::StreamExt;
756
757        let request_id = request.get("id");
758        let mut answered = request_id.is_none();
759        let mut events = response.bytes_stream().eventsource();
760        while let Some(event) = events.next().await {
761            let event = event.context("Error reading SSE response stream")?;
762            if event.data.is_empty() {
763                continue;
764            }
765            let data = match serde_json::from_str::<Value>(&event.data) {
766                Ok(msg) => {
767                    if !answered {
768                        answered = msg.get("id") == request_id
769                            && (msg.get("result").is_some() || msg.get("error").is_some());
770                    }
771                    self.observe_response(request, msg).to_string()
772                }
773                Err(_) => event.data,
774            };
775            let _ = out.send(data).await;
776        }
777        if !answered {
778            anyhow::bail!("SSE response stream ended before the response was received");
779        }
780        Ok(())
781    }
782
783    /// Resolves once a POST to the remote server has succeeded (so the session
784    /// Resolves once a legacy (2025-11-25 or earlier) session is established.
785    /// Modern servers have no standalone GET stream, so for them this never
786    /// resolves.
787    pub async fn wait_for_legacy_session(&self) {
788        let mut rx = self.legacy_session_tx.subscribe();
789        let _ = rx.wait_for(|legacy| *legacy).await;
790    }
791
792    /// Runs a re-authentication, or reuses the outcome of one that happened
793    /// while waiting.
794    ///
795    /// `observed_gen` is the credential generation used by the request that was
796    /// rejected. Requests rejected together share a single attempt: once one of
797    /// them re-authenticates, the others reuse its result (or its failure).
798    ///
799    /// For [`ReauthReason::Unauthorized`] a silent refresh is tried before the
800    /// browser login. A token from a browser login that is rejected again right
801    /// away is reported as an authentication loop instead of opening yet
802    /// another login.
803    pub async fn trigger_reauth(
804        &self,
805        observed_gen: u64,
806        metadata_url: Option<&str>,
807        scopes: Option<Vec<String>>,
808        reason: ReauthReason,
809    ) -> Result<()> {
810        let step_up = reason == ReauthReason::StepUp;
811        let attempts_before = self.reauth_state.lock().await.attempts;
812        let _guard = self.reauth_mutex.lock().await;
813
814        let mut allow_refresh = !step_up;
815        // A token from a browser login was rejected right away: only a silent
816        // refresh may be tried; another browser login would likely loop.
817        let mut refresh_only = false;
818        {
819            let state = self.reauth_state.lock().await;
820            let current_gen = self.generation();
821            let covers = |last: &LastSuccess| !step_up || last.scopes == scopes;
822
823            if current_gen != observed_gen {
824                if state.last_success.as_ref().is_some_and(covers) {
825                    info!("Credentials were renewed by another request; reusing them.");
826                    return Ok(());
827                }
828            } else if state.attempts != attempts_before {
829                // An attempt ran while we waited and failed (the generation did not move).
830                let reason = state.last_error.clone().unwrap_or_default();
831                anyhow::bail!("Re-authentication failed: {}", reason);
832            } else if let Some(last) = &state.last_success {
833                if reason != ReauthReason::Expiring
834                    && last.at.elapsed() < AUTH_LOOP_WINDOW
835                    && covers(last)
836                {
837                    if last.via_refresh {
838                        info!("The refreshed token was rejected; falling back to a new login.");
839                        allow_refresh = false;
840                    } else if allow_refresh
841                        && self.vault.get_refresh_token(&self.user_id)?.is_some()
842                    {
843                        info!("A fresh token was rejected; trying a silent refresh.");
844                        refresh_only = true;
845                    } else {
846                        error!(
847                            "Authentication loop detected: a token obtained {:?} ago was rejected.",
848                            last.at.elapsed()
849                        );
850                        anyhow::bail!(
851                            "Authentication loop detected: the server rejected a freshly issued token. \
852                             Please check your credentials and environment configuration."
853                        );
854                    }
855                }
856            }
857        }
858
859        let _ = self.suspension_tx.send(true);
860        info!("Airlock activated. Performing re-authentication...");
861
862        // `reauthenticate` bounds the interactive wait with `timeouts.auth`.
863        let result = async {
864            let auth_manager = self.ensure_auth_manager(metadata_url).await?;
865            auth_manager.enforce_issuer_binding(&self.user_id)?;
866            if allow_refresh && auth_manager.refresh(&self.user_id).await? {
867                return Ok(Some(true));
868            }
869            if refresh_only {
870                anyhow::bail!(
871                    "Authentication loop detected: the server rejected a freshly issued token. \
872                     Please check your credentials and environment configuration."
873                );
874            }
875            if reason == ReauthReason::Expiring {
876                // No refresh possible: keep the current token until it is rejected.
877                return Ok(None);
878            }
879            auth_manager
880                .reauthenticate(&self.user_id, scopes.clone(), None)
881                .await?;
882            Ok::<_, anyhow::Error>(Some(false))
883        }
884        .await;
885
886        {
887            let mut state = self.reauth_state.lock().await;
888            match &result {
889                Ok(None) => {}
890                Ok(Some(via_refresh)) => {
891                    state.attempts += 1;
892                    state.last_error = None;
893                    state.last_success = Some(LastSuccess {
894                        at: std::time::Instant::now(),
895                        scopes,
896                        via_refresh: *via_refresh,
897                    });
898                    self.reauth_count
899                        .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
900                    info!("Re-authentication successful. Deactivating Airlock...");
901                }
902                Err(e) => {
903                    state.attempts += 1;
904                    error!("Re-authentication failed: {:?}", e);
905                    state.last_error = Some(format!("{e:#}"));
906                }
907            }
908        }
909        let _ = self.suspension_tx.send(false);
910        result.map(|_| ())
911    }
912
913    async fn wait_for_airlock(&self) -> Result<()> {
914        let mut rx = self.suspension_rx.clone();
915        while *rx.borrow() {
916            rx.changed().await.context("Suspension channel closed")?;
917        }
918        Ok(())
919    }
920
921    /// Keeps the standalone GET event stream of a legacy server open,
922    /// reconnecting (with `Last-Event-ID`) and re-authenticating as needed,
923    /// and writes its messages to `stdout_tx`. Returns when the server
924    /// doesn't offer the stream (405).
925    pub async fn listen_sse(
926        &self,
927        sse_url: &str,
928        stdout_tx: tokio::sync::mpsc::Sender<String>,
929    ) -> Result<()> {
930        use futures::StreamExt;
931        use reqwest_eventsource::EventSource;
932
933        let mut loaded_gen: Option<u64> = None;
934        let mut credentials = None;
935        let mut last_event_id: Option<String> = None;
936        let mut reconnect_now = false;
937
938        loop {
939            self.wait_for_airlock().await?;
940            let gen = self.generation();
941            if loaded_gen != Some(gen) {
942                credentials = self.load_credentials()?;
943                loaded_gen = Some(gen);
944            }
945            if credentials.is_none() {
946                info!("No credentials for user in SSE listener, sending unauthenticated request to trigger discovery...");
947            }
948
949            let mut request = self
950                .build_request(reqwest::Method::GET, sse_url, credentials.as_ref(), None)
951                .await?;
952            if let Some(id) = &last_event_id {
953                // Lets the server replay what we missed (resumability).
954                request = request.header("Last-Event-ID", id);
955            }
956
957            info!("Opening SSE connection to {}...", sse_url);
958            let mut source = EventSource::new(request)?;
959
960            while let Some(event) = source.next().await {
961                match event {
962                    Ok(reqwest_eventsource::Event::Message(message)) => {
963                        tracing::debug!(event = %message.event, id = %message.id, "Received SSE message");
964                        if !message.id.is_empty() {
965                            last_event_id = Some(message.id);
966                        }
967                        if !message.data.is_empty() {
968                            let _ = stdout_tx.send(message.data).await;
969                        }
970                    }
971                    Ok(reqwest_eventsource::Event::Open) => {
972                        info!("SSE connection established");
973                        reconnect_now = false;
974                    }
975                    Err(reqwest_eventsource::Error::InvalidStatusCode(status, resp)) => {
976                        source.close();
977                        let got_nonce = self.update_rs_nonce(resp.headers());
978                        if Self::wants_dpop_nonce(status, resp.headers()) {
979                            // Retry once right away with the new nonce.
980                            reconnect_now = got_nonce && !reconnect_now;
981                            if reconnect_now {
982                                info!("SSE endpoint requires a DPoP nonce; reconnecting.");
983                            } else {
984                                error!("SSE endpoint rejected the DPoP nonce (use_dpop_nonce).");
985                            }
986                            break;
987                        }
988                        reconnect_now = false;
989                        match status {
990                            StatusCode::UNAUTHORIZED => {
991                                warn!(
992                                    "401 Unauthorized received in SSE listener ({}). Triggering re-authentication...",
993                                    sse_url
994                                );
995                                let challenge = WwwAuthenticate::parse(resp.headers());
996                                if let Err(e) = self
997                                    .reauth_for_challenge(
998                                        &challenge,
999                                        gen,
1000                                        ReauthReason::Unauthorized,
1001                                    )
1002                                    .await
1003                                {
1004                                    error!(
1005                                        "Re-authentication flow failed in SSE listener: {:?}",
1006                                        e
1007                                    );
1008                                }
1009                            }
1010                            StatusCode::METHOD_NOT_ALLOWED => {
1011                                info!(
1012                                    "Remote server does not offer an SSE stream at {} (405); SSE listener stopped.",
1013                                    sse_url
1014                                );
1015                                return Ok(());
1016                            }
1017                            _ => error!(
1018                                "SSE error: Invalid status code {} from {}. Response: {:?}",
1019                                status, sse_url, resp
1020                            ),
1021                        }
1022                        break;
1023                    }
1024                    Err(e) => {
1025                        error!("SSE error connecting to {}: {:?}", sse_url, e);
1026                        source.close();
1027                        break;
1028                    }
1029                }
1030            }
1031
1032            if reconnect_now {
1033                continue;
1034            }
1035            let t = &self.oidc_config.timeouts;
1036            let jitter_max = t.sse_retry_jitter.as_millis().max(1) as u64;
1037            let delay = t.sse_retry_base
1038                + std::time::Duration::from_millis(rand::rng().random::<u64>() % jitter_max);
1039            warn!("SSE connection lost, retrying in {:?}...", delay);
1040            tokio::time::sleep(delay).await;
1041        }
1042    }
1043}
1044
1045#[cfg(test)]
1046mod tests {
1047    use super::*;
1048    use axum::{routing::post, Router};
1049
1050    #[test]
1051    fn test_proxy_new() {
1052        let remote_url = "http://example.com/mcp";
1053        let user_id = "test_user_id";
1054        let oidc_config = OidcConfig {
1055            discovery_url: Some("http://example.com/discovery".to_string()),
1056            client_id: "test_client_id".to_string(),
1057            redirect_url: "http://localhost:8080/callback".to_string(),
1058            timeouts: crate::auth::Timeouts::fast(),
1059            ..Default::default()
1060        };
1061        let protocol_version = "2024-11-05";
1062        let auth_scheme = AuthScheme::Dpop;
1063
1064        let proxy = Proxy::new(
1065            remote_url,
1066            user_id,
1067            oidc_config.clone(),
1068            Vault::in_memory("test_service"),
1069            protocol_version,
1070            auth_scheme,
1071        );
1072
1073        assert_eq!(proxy.remote_url, remote_url);
1074        assert_eq!(proxy.user_id, user_id);
1075        assert_eq!(proxy.oidc_config.discovery_url, oidc_config.discovery_url);
1076        assert_eq!(proxy.oidc_config.client_id, oidc_config.client_id);
1077        assert_eq!(proxy.protocol_version, protocol_version);
1078        assert_eq!(proxy.auth_scheme, auth_scheme);
1079    }
1080
1081    #[test]
1082    fn test_validate_resource_metadata() {
1083        let remote_url = "http://localhost:8081/rpc";
1084        let check = |m: &str| validate_resource_metadata(Some(m), remote_url);
1085
1086        assert_eq!(
1087            check("http://localhost:8081/discovery"),
1088            Some("http://localhost:8081/discovery".to_string())
1089        );
1090        assert_eq!(check("http://attacker.com/evil"), None);
1091        assert_eq!(check("http://test.localhost:8081/discovery"), None);
1092        assert_eq!(check("http://localhost:8082/discovery"), None);
1093        assert_eq!(check("not_a_valid_url"), None);
1094        assert_eq!(validate_resource_metadata(None, remote_url), None);
1095        assert_eq!(
1096            validate_resource_metadata(Some("http://localhost:8081/d"), "not_a_valid_url"),
1097            None
1098        );
1099
1100        // Scheme downgrade and default-port equivalence.
1101        let https = "https://mcp.example.com/rpc";
1102        assert_eq!(
1103            validate_resource_metadata(Some("http://mcp.example.com/meta"), https),
1104            None
1105        );
1106        assert_eq!(
1107            validate_resource_metadata(Some("https://mcp.example.com:443/meta"), https),
1108            Some("https://mcp.example.com:443/meta".to_string())
1109        );
1110    }
1111
1112    #[tokio::test]
1113    async fn test_proxy_ensure_auth_manager_no_discovery() {
1114        let proxy = Proxy::new(
1115            "http://localhost",
1116            "user",
1117            OidcConfig {
1118                client_id: "c".into(),
1119                redirect_url: "r".into(),
1120                timeouts: crate::auth::Timeouts::fast(),
1121                ..Default::default()
1122            },
1123            Vault::in_memory("svc"),
1124            "v1",
1125            AuthScheme::Bearer,
1126        );
1127        // This should fail because no discovery and no overrides
1128        let res = proxy.ensure_auth_manager(None).await;
1129        assert!(res.is_err());
1130    }
1131
1132    #[tokio::test]
1133    async fn test_proxy_handle_request_no_content() -> Result<()> {
1134        let mcp_app = Router::new().route(
1135            "/rpc",
1136            post(|| async move { axum::http::StatusCode::NO_CONTENT }),
1137        );
1138        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
1139        let addr = listener.local_addr()?;
1140        let rpc_url = format!("http://127.0.0.1:{}/rpc", addr.port());
1141        tokio::spawn(async move {
1142            let _ = axum::serve(listener, mcp_app).await;
1143        });
1144
1145        let vault = Vault::in_memory("svc");
1146        let proxy = Proxy::new(
1147            &rpc_url,
1148            "user",
1149            OidcConfig {
1150                client_id: "c".into(),
1151                redirect_url: "r".into(),
1152                timeouts: crate::auth::Timeouts::fast(),
1153                ..Default::default()
1154            },
1155            vault.clone(),
1156            "v1",
1157            AuthScheme::Bearer,
1158        );
1159        vault.store_token("user", "token")?;
1160        vault.store_dpop_key("user", &crate::crypto::DpopKey::generate().to_bytes())?;
1161
1162        let res = proxy
1163            .call(serde_json::json!({"jsonrpc": "2.0", "id": 1, "method": "test"}))
1164            .await?;
1165        assert_eq!(res, None);
1166        Ok(())
1167    }
1168
1169    fn test_proxy(url: &str, oidc: OidcConfig) -> Arc<Proxy> {
1170        Proxy::new(
1171            url,
1172            "user",
1173            OidcConfig {
1174                client_id: "c".into(),
1175                timeouts: crate::auth::Timeouts::fast(),
1176                ..oidc
1177            },
1178            Vault::in_memory("svc"),
1179            "v1",
1180            AuthScheme::Bearer,
1181        )
1182    }
1183
1184    /// A proxy whose AuthManager can never be built (no discovery possible).
1185    fn undiscoverable_proxy() -> Arc<Proxy> {
1186        test_proxy(
1187            "http://localhost:1/rpc",
1188            OidcConfig {
1189                redirect_url: "r".into(),
1190                ..Default::default()
1191            },
1192        )
1193    }
1194
1195    #[tokio::test]
1196    async fn test_proxy_handle_request_max_retries() -> Result<()> {
1197        use std::sync::atomic::{AtomicUsize, Ordering};
1198        let counter = Arc::new(AtomicUsize::new(0));
1199
1200        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
1201        let rpc_url = format!("http://{}/rpc", listener.local_addr()?);
1202        let proxy = test_proxy(
1203            &rpc_url,
1204            OidcConfig {
1205                redirect_url: "http://127.0.0.1:1/callback".into(),
1206                ..Default::default()
1207            },
1208        );
1209
1210        // The server rejects every request, and "someone" renews the credentials
1211        // each time, so every re-auth is skipped and the request is retried.
1212        let (c, p) = (counter.clone(), proxy.clone());
1213        let mcp_app = Router::new().route(
1214            "/rpc",
1215            post(move || {
1216                let (c, p) = (c.clone(), p.clone());
1217                async move {
1218                    c.fetch_add(1, Ordering::SeqCst);
1219                    {
1220                        let mut state = p.reauth_state.lock().await;
1221                        state.last_success = Some(LastSuccess {
1222                            at: std::time::Instant::now(),
1223                            scopes: None,
1224                            via_refresh: false,
1225                        });
1226                    }
1227                    p.reauth_count.fetch_add(1, Ordering::SeqCst);
1228                    (axum::http::StatusCode::UNAUTHORIZED, "Unauthorized")
1229                }
1230            }),
1231        );
1232        tokio::spawn(async move {
1233            let _ = axum::serve(listener, mcp_app).await;
1234        });
1235
1236        let err = proxy
1237            .call(serde_json::json!({"jsonrpc": "2.0", "id": 1, "method": "test"}))
1238            .await
1239            .unwrap_err();
1240        assert!(err.to_string().contains("Maximum retry attempts reached"));
1241        // Initial attempt + 2 retries.
1242        assert_eq!(counter.load(Ordering::SeqCst), 3);
1243        Ok(())
1244    }
1245
1246    #[tokio::test]
1247    async fn test_proxy_reauth_failure_resets_circuit_breaker() -> Result<()> {
1248        let proxy = undiscoverable_proxy();
1249
1250        let first = proxy
1251            .trigger_reauth(0, None, None, ReauthReason::Unauthorized)
1252            .await;
1253        assert!(first.is_err());
1254        assert!(proxy.reauth_state.lock().await.last_success.is_none());
1255        assert!(!*proxy.suspension_rx.borrow(), "airlock must be released");
1256
1257        // A failed attempt must not arm the loop detector: retrying later runs
1258        // the flow again instead of reporting an authentication loop.
1259        let second = proxy
1260            .trigger_reauth(0, None, None, ReauthReason::Unauthorized)
1261            .await;
1262        let err = second.unwrap_err().to_string();
1263        assert!(!err.contains("Authentication loop detected"), "{err}");
1264        assert_eq!(proxy.reauth_state.lock().await.attempts, 2);
1265        Ok(())
1266    }
1267
1268    #[tokio::test]
1269    async fn test_proxy_trigger_reauth_reuses_renewal_by_other_request() -> Result<()> {
1270        let proxy = undiscoverable_proxy();
1271        // Another request re-authenticated after our credentials were loaded.
1272        proxy.reauth_state.lock().await.last_success = Some(LastSuccess {
1273            at: std::time::Instant::now(),
1274            scopes: None,
1275            via_refresh: false,
1276        });
1277        proxy
1278            .reauth_count
1279            .store(1, std::sync::atomic::Ordering::SeqCst);
1280
1281        // No network access happens: the renewal is reused.
1282        proxy
1283            .trigger_reauth(0, None, None, ReauthReason::Unauthorized)
1284            .await?;
1285        assert_eq!(proxy.reauth_state.lock().await.attempts, 0);
1286        Ok(())
1287    }
1288
1289    #[tokio::test]
1290    async fn test_proxy_trigger_reauth_step_up_not_covered_by_plain_renewal() -> Result<()> {
1291        let proxy = undiscoverable_proxy();
1292        proxy.reauth_state.lock().await.last_success = Some(LastSuccess {
1293            at: std::time::Instant::now(),
1294            scopes: None,
1295            via_refresh: false,
1296        });
1297        proxy
1298            .reauth_count
1299            .store(1, std::sync::atomic::Ordering::SeqCst);
1300
1301        // A step-up for new scopes still needs its own login (which fails here).
1302        let res = proxy
1303            .trigger_reauth(0, None, Some(vec!["admin".into()]), ReauthReason::StepUp)
1304            .await;
1305        assert!(res.is_err());
1306        assert_eq!(proxy.reauth_state.lock().await.attempts, 1);
1307        Ok(())
1308    }
1309
1310    #[tokio::test]
1311    async fn test_proxy_trigger_reauth_detects_loop() -> Result<()> {
1312        let proxy = undiscoverable_proxy();
1313        // We re-authenticated just now and the request with that token (gen 1)
1314        // was rejected again.
1315        proxy.reauth_state.lock().await.last_success = Some(LastSuccess {
1316            at: std::time::Instant::now(),
1317            scopes: None,
1318            via_refresh: false,
1319        });
1320        proxy
1321            .reauth_count
1322            .store(1, std::sync::atomic::Ordering::SeqCst);
1323
1324        let err = proxy
1325            .trigger_reauth(1, None, None, ReauthReason::Unauthorized)
1326            .await
1327            .unwrap_err();
1328        assert!(err.to_string().contains("Authentication loop detected"));
1329
1330        // With a refresh token, a silent refresh is tried first (which fails
1331        // here for lack of discovery, with a different error).
1332        proxy.vault.store_refresh_token("user", "r")?;
1333        let err = proxy
1334            .trigger_reauth(1, None, None, ReauthReason::Unauthorized)
1335            .await
1336            .unwrap_err();
1337        assert!(
1338            !err.to_string().contains("Authentication loop detected"),
1339            "{err}"
1340        );
1341        proxy.vault.delete_refresh_token("user")?;
1342
1343        // A rejected *refreshed* token is not a loop: it falls back to a login
1344        // (which fails here for lack of discovery, with a different error).
1345        proxy.reauth_state.lock().await.last_success = Some(LastSuccess {
1346            at: std::time::Instant::now(),
1347            scopes: None,
1348            via_refresh: true,
1349        });
1350        let err = proxy
1351            .trigger_reauth(1, None, None, ReauthReason::Unauthorized)
1352            .await
1353            .unwrap_err();
1354        assert!(!err.to_string().contains("Authentication loop detected"));
1355
1356        // An old success does not count as a loop.
1357        proxy.reauth_state.lock().await.last_success = Some(LastSuccess {
1358            at: std::time::Instant::now() - AUTH_LOOP_WINDOW * 2,
1359            scopes: None,
1360            via_refresh: false,
1361        });
1362        let err = proxy
1363            .trigger_reauth(1, None, None, ReauthReason::Unauthorized)
1364            .await
1365            .unwrap_err();
1366        assert!(!err.to_string().contains("Authentication loop detected"));
1367        Ok(())
1368    }
1369
1370    #[tokio::test]
1371    async fn test_proxy_concurrent_failures_share_one_attempt() -> Result<()> {
1372        use std::sync::atomic::{AtomicUsize, Ordering};
1373        // Discovery is slow and fails, so the second request waits on the
1374        // first attempt and must reuse its failure instead of starting a new one.
1375        let hits = Arc::new(AtomicUsize::new(0));
1376        let h = hits.clone();
1377        let app = Router::new().fallback(move || {
1378            let h = h.clone();
1379            async move {
1380                h.fetch_add(1, Ordering::SeqCst);
1381                tokio::time::sleep(std::time::Duration::from_millis(100)).await;
1382                axum::http::StatusCode::NOT_FOUND
1383            }
1384        });
1385        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
1386        let base = format!("http://{}", listener.local_addr()?);
1387        tokio::spawn(async move {
1388            let _ = axum::serve(listener, app).await;
1389        });
1390        let proxy = test_proxy(
1391            &format!("{base}/rpc"),
1392            OidcConfig {
1393                redirect_url: "r".into(),
1394                ..Default::default()
1395            },
1396        );
1397
1398        let p1 = proxy.clone();
1399        let first = tokio::spawn(async move {
1400            p1.trigger_reauth(0, None, None, ReauthReason::Unauthorized)
1401                .await
1402        });
1403        tokio::time::sleep(std::time::Duration::from_millis(30)).await;
1404        let second = proxy
1405            .trigger_reauth(0, None, None, ReauthReason::Unauthorized)
1406            .await;
1407
1408        assert!(first.await?.is_err());
1409        let err = second.unwrap_err().to_string();
1410        assert!(err.contains("Re-authentication failed"), "{err}");
1411        assert_eq!(proxy.reauth_state.lock().await.attempts, 1);
1412        // One discovery run: the path-inserted and the root well-known URIs.
1413        assert_eq!(hits.load(Ordering::SeqCst), 2);
1414        Ok(())
1415    }
1416
1417    #[tokio::test]
1418    async fn test_proxy_wait_for_airlock() -> Result<()> {
1419        let proxy = Proxy::new(
1420            "http://localhost:1/rpc",
1421            "user",
1422            OidcConfig {
1423                client_id: "c".into(),
1424                redirect_url: "r".into(),
1425                timeouts: crate::auth::Timeouts::fast(),
1426                ..Default::default()
1427            },
1428            Vault::in_memory("svc"),
1429            "v1",
1430            AuthScheme::Bearer,
1431        );
1432
1433        // Initially not suspended
1434        let p = proxy.clone();
1435        let _ = tokio::time::timeout(std::time::Duration::from_millis(100), p.wait_for_airlock())
1436            .await?;
1437
1438        // Manually activate airlock
1439        let _ = proxy.suspension_tx.send(true);
1440        let p2 = proxy.clone();
1441        let handle = tokio::spawn(async move {
1442            let _ = p2.wait_for_airlock().await;
1443        });
1444
1445        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1446        let _ = proxy.suspension_tx.send(false);
1447
1448        tokio::time::timeout(std::time::Duration::from_millis(100), handle).await??;
1449        Ok(())
1450    }
1451
1452    #[tokio::test]
1453    async fn test_listen_sse_retry_logic() -> Result<()> {
1454        let (tx, _rx) = tokio::sync::mpsc::channel(1);
1455        let proxy = Proxy::new(
1456            "http://localhost:1/rpc",
1457            "user",
1458            OidcConfig {
1459                client_id: "c".into(),
1460                redirect_url: "r".into(),
1461                timeouts: crate::auth::Timeouts::fast(),
1462                ..Default::default()
1463            },
1464            Vault::in_memory("svc"),
1465            "v1",
1466            AuthScheme::Bearer,
1467        );
1468
1469        let p = proxy.clone();
1470        let handle = tokio::spawn(async move {
1471            let _ = p.listen_sse("http://localhost:1/sse", tx).await;
1472        });
1473
1474        tokio::time::sleep(std::time::Duration::from_millis(100)).await;
1475        handle.abort();
1476        Ok(())
1477    }
1478
1479    #[tokio::test]
1480    async fn test_listen_sse_failure() -> Result<()> {
1481        let (tx, _rx) = tokio::sync::mpsc::channel(1);
1482        let proxy = Proxy::new(
1483            "http://localhost:1/rpc",
1484            "user",
1485            OidcConfig {
1486                client_id: "c".into(),
1487                redirect_url: "r".into(),
1488                timeouts: crate::auth::Timeouts::fast(),
1489                ..Default::default()
1490            },
1491            Vault::in_memory("svc"),
1492            "v1",
1493            AuthScheme::Bearer,
1494        );
1495        // It will retry infinitely, so we just want to see it starting and failing once.
1496        let res = tokio::time::timeout(
1497            std::time::Duration::from_millis(100),
1498            proxy.listen_sse("http://localhost:1/sse", tx),
1499        )
1500        .await;
1501        assert!(res.is_err()); // Timeout means it's still retrying
1502        Ok(())
1503    }
1504}