1use 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
24pub struct Proxy {
26 http_client: Client,
28 remote_url: String,
30 suspension_rx: watch::Receiver<bool>,
32 suspension_tx: watch::Sender<bool>,
34 vault: Vault,
36 user_id: String,
38 oidc_config: OidcConfig,
40 pub auth_manager: Arc<RwLock<Option<Arc<AuthManager>>>>,
42 protocol_version: String,
44 auth_scheme: AuthScheme,
46 session_id: Mutex<Option<String>>,
48 reauth_mutex: Mutex<()>,
50 reauth_state: Mutex<ReauthState>,
52 reauth_count: Arc<std::sync::atomic::AtomicU64>,
54 legacy_session_tx: watch::Sender<bool>,
57 rs_nonce: std::sync::Mutex<Option<String>>,
59 proactive_refresh_gen: std::sync::Mutex<Option<u64>>,
61 negotiated_version: std::sync::Mutex<Option<String>>,
63 tool_headers: std::sync::Mutex<std::collections::HashMap<String, mcp::ToolHeaders>>,
65}
66
67const AUTH_LOOP_WINDOW: std::time::Duration = std::time::Duration::from_secs(5);
70
71#[derive(Clone, Copy, Debug, PartialEq, Eq)]
73pub enum ReauthReason {
74 Unauthorized,
76 StepUp,
78 Expiring,
80}
81
82#[derive(Default)]
84struct ReauthState {
85 attempts: u64,
87 last_error: Option<String>,
89 last_success: Option<LastSuccess>,
91}
92
93struct LastSuccess {
94 at: std::time::Instant,
95 scopes: Option<Vec<String>>,
97 via_refresh: bool,
99}
100
101enum Outcome {
103 Done,
105 HeaderMismatch(Value),
107}
108
109const EXPIRY_MARGIN_SECS: u64 = 30;
111
112struct Credentials {
114 token: String,
115 key: DpopKey,
116 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 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 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 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 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 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 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 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 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(), 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 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 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 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 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 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 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 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 async fn reauth_for_challenge(
497 &self,
498 challenge: &WwwAuthenticate,
499 observed_gen: u64,
500 reason: ReauthReason,
501 ) -> Result<()> {
502 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 let p = proxy.clone();
1435 let _ = tokio::time::timeout(std::time::Duration::from_millis(100), p.wait_for_airlock())
1436 .await?;
1437
1438 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 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()); Ok(())
1503 }
1504}