1use crate::crypto::DpopKey;
13use crate::discovery;
14use crate::net;
15use crate::vault::{CredentialMeta, Vault};
16use crate::Result;
17use anyhow::Context;
18use axum::{
19 extract::{Query, State},
20 response::IntoResponse,
21 routing::get,
22 Router,
23};
24use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
25use colored::Colorize;
26use rand_core::{OsRng, RngCore};
27use reqwest::Client as HttpClient;
28use serde::Deserialize;
29use sha2::{Digest, Sha256};
30use std::net::SocketAddr;
31use std::sync::Arc;
32use std::time::Duration;
33use tokio::sync::oneshot;
34use tracing::{error, info, warn};
35
36#[derive(Clone)]
38pub struct OidcConfig {
39 pub discovery_url: Option<String>,
41 pub client_id: String,
43 pub redirect_url: String,
45 pub auth_url_override: Option<String>,
47 pub token_url_override: Option<String>,
49 pub par_url_override: Option<String>,
51 pub internal_url_tx: Arc<tokio::sync::Mutex<Option<oneshot::Sender<String>>>>,
53 pub internal_callback_tx: Arc<tokio::sync::Mutex<Option<oneshot::Sender<SocketAddr>>>>,
55 pub template_dir: Option<std::path::PathBuf>,
57 pub timeouts: Timeouts,
59 pub allow_insecure_http: bool,
61 pub expected_issuer: Option<String>,
64 pub request_offline_access: bool,
68}
69
70impl Default for OidcConfig {
71 fn default() -> Self {
72 Self {
73 discovery_url: None,
74 client_id: String::new(),
75 redirect_url: String::new(),
76 auth_url_override: None,
77 token_url_override: None,
78 par_url_override: None,
79 internal_url_tx: Arc::new(tokio::sync::Mutex::new(None)),
80 internal_callback_tx: Arc::new(tokio::sync::Mutex::new(None)),
81 template_dir: None,
82 timeouts: Timeouts::default(),
83 allow_insecure_http: false,
84 expected_issuer: None,
85 request_offline_access: false,
86 }
87 }
88}
89
90#[derive(Clone, Debug, PartialEq, Eq)]
92pub struct Timeouts {
93 pub auth: Duration,
95 pub bind_retry: Duration,
97 pub sse_retry_base: Duration,
99 pub sse_retry_jitter: Duration,
101}
102
103impl Default for Timeouts {
104 fn default() -> Self {
105 Self {
106 auth: Duration::from_secs(300),
107 bind_retry: Duration::from_secs(1),
108 sse_retry_base: Duration::from_secs(5),
109 sse_retry_jitter: Duration::from_secs(2),
110 }
111 }
112}
113
114impl Timeouts {
115 pub fn fast() -> Self {
117 Self {
118 auth: Duration::from_millis(500),
119 bind_retry: Duration::from_millis(10),
120 sse_retry_base: Duration::from_millis(10),
121 sse_retry_jitter: Duration::from_millis(10),
122 }
123 }
124}
125
126pub(crate) const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
128const AS_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
130
131impl std::fmt::Debug for OidcConfig {
132 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
133 f.debug_struct("OidcConfig")
134 .field("discovery_url", &self.discovery_url)
135 .field("client_id", &self.client_id)
136 .field("redirect_url", &self.redirect_url)
137 .field("template_dir", &self.template_dir)
138 .field("timeouts", &self.timeouts)
139 .field("allow_insecure_http", &self.allow_insecure_http)
140 .field("expected_issuer", &self.expected_issuer)
141 .finish()
142 }
143}
144
145#[derive(Clone)]
147pub struct AuthManager {
148 client_id: String,
150 auth_url: String,
152 token_url: String,
154 par_url: String,
156 redirect_url: String,
158 resource: String,
160 http_client: HttpClient,
162 vault: Vault,
164 internal_url_tx: Arc<tokio::sync::Mutex<Option<oneshot::Sender<String>>>>,
166 internal_callback_tx: Arc<tokio::sync::Mutex<Option<oneshot::Sender<SocketAddr>>>>,
168 success_html: Arc<String>,
170 failure_html: Arc<String>,
172 issuer_name: String,
174 resource_name: String,
176 timeouts: Timeouts,
178 issuer: Option<String>,
180 iss_required: bool,
182 as_nonce: Arc<std::sync::Mutex<Option<String>>>,
184 as_scopes: Option<Vec<String>>,
186 resource_scopes: Option<Vec<String>>,
188 request_offline_access: bool,
190}
191
192#[derive(Deserialize, Default)]
194struct AuthCallback {
195 code: Option<String>,
196 state: Option<String>,
197 error: Option<String>,
198 error_description: Option<String>,
199 iss: Option<String>,
200}
201
202#[derive(Deserialize)]
203struct ParResponse {
204 request_uri: String,
205}
206
207#[derive(Deserialize)]
209struct TokenResponse {
210 access_token: String,
211 token_type: Option<String>,
212 refresh_token: Option<String>,
213 expires_in: Option<u64>,
214}
215
216impl AuthManager {
217 pub async fn discover(
222 mut oidc_config: OidcConfig,
223 resource: String,
224 vault: Vault,
225 resource_metadata_url: Option<&str>,
226 ) -> Result<Self> {
227 let http_client = HttpClient::builder()
228 .connect_timeout(CONNECT_TIMEOUT)
229 .timeout(AS_REQUEST_TIMEOUT)
230 .build()
231 .context("Failed to build HTTP client")?;
232
233 let overrides_complete = oidc_config.auth_url_override.is_some()
234 && oidc_config.token_url_override.is_some()
235 && oidc_config.par_url_override.is_some();
236
237 let (metadata, mut resource_name, resource_scopes) = if overrides_complete {
240 (None, None, None)
241 } else if let Some(url) = oidc_config.discovery_url.as_deref() {
242 info!("Using configured OIDC discovery URL {}", url);
243 net::require_secure_url(url, "OIDC discovery URL", oidc_config.allow_insecure_http)?;
244 let metadata = discovery::fetch_configured_metadata(&http_client, url).await?;
245 (Some(metadata), None, None)
246 } else {
247 info!("Discovering the authorization server for {}...", resource);
248 let found = discovery::discover_from_resource(
249 &http_client,
250 &resource,
251 resource_metadata_url,
252 oidc_config.allow_insecure_http,
253 )
254 .await?;
255 (
256 Some(found.metadata),
257 found.resource_name,
258 found.resource_scopes,
259 )
260 };
261
262 if let Some(m) = &metadata {
263 if let Some(expected) = &oidc_config.expected_issuer {
264 if &m.issuer != expected {
265 anyhow::bail!(
266 "The authorization server '{}' is not the one this client is registered \
267 with (--oidc-issuer '{}')",
268 m.issuer,
269 expected
270 );
271 }
272 }
273 if !m.supports_pkce_s256() {
275 anyhow::bail!(
276 "Authorization server '{}' does not advertise PKCE S256 in \
277 code_challenge_methods_supported; refusing to proceed",
278 m.issuer
279 );
280 }
281 if is_url_client_id(&oidc_config.client_id) && !m.client_id_metadata_document_supported
282 {
283 warn!(
284 "Client ID '{}' is a metadata document URL, but '{}' does not advertise \
285 client_id_metadata_document_supported.",
286 oidc_config.client_id, m.issuer
287 );
288 }
289 }
290 if is_url_client_id(&oidc_config.client_id) {
291 check_cimd_client_id(&oidc_config.client_id)?;
292 }
293
294 let auth_url = oidc_config
295 .auth_url_override
296 .take()
297 .or_else(|| metadata.as_ref().map(|m| m.authorization_endpoint.clone()))
298 .context("No authorization endpoint (set --kc-auth-url or a discovery URL)")?;
299 let token_url = oidc_config
300 .token_url_override
301 .take()
302 .or_else(|| metadata.as_ref().map(|m| m.token_endpoint.clone()))
303 .context("No token endpoint (set --kc-token-url or a discovery URL)")?;
304 let par_url = oidc_config
305 .par_url_override
306 .take()
307 .or_else(|| {
308 metadata
309 .as_ref()
310 .and_then(|m| m.pushed_authorization_request_endpoint.clone())
311 })
312 .context(
313 "Authorization server metadata has no pushed_authorization_request_endpoint \
314 and no --kc-par-url override was provided",
315 )?;
316 for (url, what) in [
317 (&auth_url, "authorization endpoint"),
318 (&token_url, "token endpoint"),
319 (&par_url, "PAR endpoint"),
320 ] {
321 net::require_secure_url(url, what, oidc_config.allow_insecure_http)?;
322 }
323
324 let issuer = metadata
325 .as_ref()
326 .map(|m| m.issuer.clone())
327 .or_else(|| oidc_config.expected_issuer.clone());
328 let as_scopes = metadata.as_ref().and_then(|m| m.scopes_supported.clone());
329 let iss_required = metadata
330 .as_ref()
331 .is_some_and(|m| m.authorization_response_iss_parameter_supported);
332 let issuer_name = metadata
333 .map(|m| m.organization_name.unwrap_or(m.issuer))
334 .unwrap_or_else(|| "Custom Provider".to_string());
335
336 if resource_name.is_none() {
337 resource_name = discovery::fetch_resource_name(&http_client, &resource).await;
338 }
339 let resource_name = resource_name.unwrap_or_else(|| resource.clone());
340
341 let (success_html, failure_html) = if let Some(dir) = &oidc_config.template_dir {
342 let (success_res, failure_res) = tokio::join!(
343 tokio::fs::read_to_string(dir.join("success.html")),
344 tokio::fs::read_to_string(dir.join("failure.html"))
345 );
346 (
347 success_res.unwrap_or_else(|_| crate::templates::DEFAULT_SUCCESS_HTML.to_string()),
348 failure_res.unwrap_or_else(|_| crate::templates::DEFAULT_FAILURE_HTML.to_string()),
349 )
350 } else {
351 (
352 crate::templates::DEFAULT_SUCCESS_HTML.to_string(),
353 crate::templates::DEFAULT_FAILURE_HTML.to_string(),
354 )
355 };
356
357 Ok(Self {
358 client_id: oidc_config.client_id,
359 auth_url,
360 token_url,
361 par_url,
362 redirect_url: oidc_config.redirect_url,
363 resource,
364 http_client,
365 vault,
366 internal_url_tx: oidc_config.internal_url_tx,
367 internal_callback_tx: oidc_config.internal_callback_tx,
368 success_html: Arc::new(success_html),
369 failure_html: Arc::new(failure_html),
370 issuer_name,
371 resource_name,
372 timeouts: oidc_config.timeouts,
373 issuer,
374 iss_required,
375 as_nonce: Arc::default(),
376 as_scopes,
377 resource_scopes,
378 request_offline_access: oidc_config.request_offline_access,
379 })
380 }
381
382 pub async fn set_internal_url_tx(&self, tx: oneshot::Sender<String>) {
385 let mut lock = self.internal_url_tx.lock().await;
386 *lock = Some(tx);
387 }
388
389 pub async fn set_internal_callback_tx(&self, tx: oneshot::Sender<SocketAddr>) {
391 let mut lock = self.internal_callback_tx.lock().await;
392 *lock = Some(tx);
393 }
394
395 pub async fn reauthenticate(
397 &self,
398 user_id: &str,
399 scopes: Option<Vec<String>>,
400 url_tx: Option<oneshot::Sender<String>>,
401 ) -> Result<()> {
402 info!(
403 "Starting FAPI 2.0 re-authentication flow for user '{}'...",
404 user_id
405 );
406
407 let dpop_key = DpopKey::generate();
410
411 let pkce_verifier = random_urlsafe(32);
413 let pkce_challenge = pkce_s256_challenge(&pkce_verifier);
414 let state_val = random_urlsafe(16);
415
416 let expected_state = state_val.clone();
418 let (server, rx) = self.setup_loopback_server(expected_state).await?;
419
420 let previous = self
422 .bound_meta(user_id)?
423 .map(|m| m.scopes)
424 .unwrap_or_default();
425 let scopes = self.select_scopes(scopes, &previous);
426 let dpop_jkt = dpop_key.jkt()?;
427 let par_data = match self
428 .perform_par_request(&pkce_challenge, &state_val, &scopes, &dpop_jkt)
429 .await
430 {
431 Ok(data) => data,
432 Err(e) => {
433 server.stop().await;
434 return Err(e);
435 }
436 };
437
438 if let Err(e) = self.open_auth_url(&par_data, url_tx).await {
440 server.stop().await;
441 return Err(e);
442 }
443
444 let callback = tokio::time::timeout(self.timeouts.auth, rx).await;
446 server.stop().await;
447 let code = match callback {
448 Ok(Ok(Ok(code))) => code,
449 Ok(Ok(Err(reason))) => anyhow::bail!("Authorization failed: {}", reason),
450 _ => anyhow::bail!("Authentication timed out or failed to receive callback"),
451 };
452
453 info!("Step 2: Exchanging code for DPoP-bound token...");
455 self.manual_token_exchange(user_id, &code, &pkce_verifier, &dpop_key, &scopes)
456 .await?;
457
458 Ok(())
459 }
460
461 async fn setup_loopback_server(
462 &self,
463 expected_state: String,
464 ) -> Result<(LoopbackServer, oneshot::Receiver<CallbackResult>)> {
465 let redirect = net::require_loopback_redirect(&self.redirect_url)?;
466 let (tx, rx) = oneshot::channel::<CallbackResult>();
467 let tx = Arc::new(tokio::sync::Mutex::new(Some(tx)));
468
469 let app = Router::new()
470 .route("/callback", get(handle_callback))
471 .layer(axum::middleware::map_response(
474 |mut response: axum::response::Response| async move {
475 response.headers_mut().insert(
476 axum::http::header::CONNECTION,
477 axum::http::HeaderValue::from_static("close"),
478 );
479 response
480 },
481 ))
482 .with_state(AuthServerState {
483 expected_state,
484 tx,
485 success_html: self.success_html.clone(),
486 failure_html: self.failure_html.clone(),
487 issuer_name: self.issuer_name.clone(),
488 resource_name: self.resource_name.clone(),
489 expected_issuer: self.issuer.clone(),
490 iss_required: self.iss_required,
491 });
492
493 let addr: SocketAddr = redirect
494 .socket_addrs(|| None)?
495 .first()
496 .copied()
497 .context("Failed to parse redirect URL into socket address")?;
498
499 let mut listener = None;
500 for i in 0..5 {
501 match tokio::net::TcpListener::bind(addr).await {
502 Ok(l) => {
503 listener = Some(l);
504 break;
505 }
506 Err(e) if e.kind() == std::io::ErrorKind::AddrInUse => {
507 if i == 4 {
508 return Err(anyhow::anyhow!(e).context(format!("Failed to bind to {} after 5 retries. Someone else is using this port.", addr)));
509 }
510 warn!(
511 "Address {} already in use, retrying... (attempt {})",
512 addr,
513 i + 1
514 );
515 tokio::time::sleep(self.timeouts.bind_retry).await;
516 }
517 Err(e) => return Err(e.into()),
518 }
519 }
520
521 let listener = listener.context("Failed to bind loopback listener after retries")?;
522 let local_addr = listener.local_addr()?;
523
524 {
526 let mut lock = self.internal_callback_tx.lock().await;
527 if let Some(tx_addr) = lock.take() {
528 let _ = tx_addr.send(local_addr);
529 }
530 }
531
532 let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>();
533 let handle = tokio::spawn(async move {
534 let serve = axum::serve(listener, app).with_graceful_shutdown(async {
535 let _ = shutdown_rx.await;
536 });
537 if let Err(e) = serve.await {
538 error!("Loopback server error: {:?}", e);
539 }
540 });
541
542 Ok((
543 LoopbackServer {
544 handle,
545 shutdown: Some(shutdown_tx),
546 },
547 rx,
548 ))
549 }
550
551 async fn perform_par_request(
552 &self,
553 pkce_challenge: &str,
554 state_val: &str,
555 scopes: &[String],
556 dpop_jkt: &str,
557 ) -> Result<ParResponse> {
558 info!("Step 1: Pushed Authorization Request (PAR)...");
559 let mut par_params = vec![
560 ("client_id", self.client_id.as_str()),
561 ("response_type", "code"),
562 ("redirect_uri", self.redirect_url.as_str()),
563 ("code_challenge", pkce_challenge),
564 ("code_challenge_method", "S256"),
565 ("state", state_val),
566 ("resource", self.resource.as_str()),
567 ("dpop_jkt", dpop_jkt),
569 ];
570
571 let scope_str = scopes.join(" ");
572 if !scope_str.is_empty() {
573 par_params.push(("scope", &scope_str));
574 }
575
576 let par_res = self
577 .http_client
578 .post(&self.par_url)
579 .form(&par_params)
580 .send()
581 .await?;
582
583 if !par_res.status().is_success() {
584 let error_text = par_res.text().await?;
585 error!("PAR request failed: {}", error_text);
586 anyhow::bail!("PAR request failed: {}", error_text);
587 }
588
589 let par_data: ParResponse = par_res.json().await?;
590 Ok(par_data)
591 }
592
593 async fn open_auth_url(
594 &self,
595 par_data: &ParResponse,
596 url_tx: Option<oneshot::Sender<String>>,
597 ) -> Result<()> {
598 let auth_url = build_authorize_url(&self.auth_url, &self.client_id, &par_data.request_uri)?;
599
600 eprintln!(
601 "{}",
602 "****************************************************************".yellow()
603 );
604 eprintln!(
605 "{}",
606 "🔐 ACTION REQUIRED: Please visit the following URL to authenticate:"
607 .bold()
608 .yellow()
609 );
610 eprintln!("{}", auth_url.bold().cyan());
611 eprintln!(
612 "{}",
613 "****************************************************************".yellow()
614 );
615
616 let skip_open = std::env::var("MCP_AIRLOCK_SKIP_OPEN_BROWSER").is_ok();
618 let mut has_listener = url_tx.is_some();
619
620 if !has_listener {
621 let lock = self.internal_url_tx.lock().await;
622 has_listener = lock.is_some();
623 }
624
625 if !skip_open && !has_listener {
626 let is_safe_url = url::Url::parse(&auth_url)
627 .is_ok_and(|u| u.scheme() == "http" || u.scheme() == "https");
628
629 if is_safe_url {
630 if let Err(e) = open::that(&auth_url) {
631 warn!(
632 "Failed to open browser automatically: {}. Please copy the URL above.",
633 e
634 );
635 }
636 } else {
637 warn!("Skipping automatic browser open: URL scheme is not http or https. Please copy the URL above.");
638 }
639 } else {
640 info!("Skipping automatic browser open (internal listener or skip flag present).");
641 }
642
643 if let Some(tx_url) = url_tx {
645 let _ = tx_url.send(auth_url.clone());
646 } else {
647 let mut lock = self.internal_url_tx.lock().await;
648 if let Some(tx_url) = lock.take() {
649 let _ = tx_url.send(auth_url.clone());
650 }
651 }
652 Ok(())
653 }
654
655 async fn manual_token_exchange(
658 &self,
659 user_id: &str,
660 code: &str,
661 pkce_verifier: &str,
662 dpop_key: &DpopKey,
663 scopes: &[String],
664 ) -> Result<()> {
665 let params = [
666 ("grant_type", "authorization_code"),
667 ("client_id", self.client_id.as_str()),
668 ("code", code),
669 ("redirect_uri", self.redirect_url.as_str()),
670 ("code_verifier", pkce_verifier),
671 ("resource", self.resource.as_str()),
672 ];
673 let tokens = self
674 .token_request(¶ms, dpop_key)
675 .await?
676 .map_err(|err| anyhow::anyhow!("Token exchange failed: {}", err))?;
677
678 self.vault.store_dpop_key(user_id, &dpop_key.to_bytes())?;
679 self.store_tokens(user_id, &tokens, Some(scopes))?;
680 if tokens.refresh_token.is_none() {
681 self.vault.delete_refresh_token(user_id)?;
683 }
684 info!("Successfully acquired and stored DPoP-bound token.");
685 Ok(())
686 }
687
688 pub async fn refresh(&self, user_id: &str) -> Result<bool> {
695 let Some(refresh_token) = self.vault.get_refresh_token(user_id)? else {
696 return Ok(false);
697 };
698 if self.vault.get_meta(user_id)?.is_some() && self.bound_meta(user_id)?.is_none() {
700 self.vault.delete_refresh_token(user_id)?;
701 return Ok(false);
702 }
703 let Some(key_bytes) = self.vault.get_dpop_key(user_id)? else {
705 self.vault.delete_refresh_token(user_id)?;
706 return Ok(false);
707 };
708 let dpop_key = DpopKey::from_bytes(&key_bytes)?;
709
710 info!("Refreshing the access token...");
711 let params = [
712 ("grant_type", "refresh_token"),
713 ("client_id", self.client_id.as_str()),
714 ("refresh_token", refresh_token.as_str()),
715 ("resource", self.resource.as_str()),
716 ];
717 match self.token_request(¶ms, &dpop_key).await? {
718 Ok(tokens) => {
719 self.store_tokens(user_id, &tokens, None)?;
720 info!("Access token refreshed.");
721 Ok(true)
722 }
723 Err(err) => {
724 warn!("Refresh token rejected ({}); a new login is needed.", err);
725 self.vault.delete_refresh_token(user_id)?;
726 Ok(false)
727 }
728 }
729 }
730
731 async fn token_request(
737 &self,
738 params: &[(&str, &str)],
739 dpop_key: &DpopKey,
740 ) -> Result<std::result::Result<TokenResponse, String>> {
741 let mut retried = false;
742 let res = loop {
743 let nonce = self.as_nonce.lock().ok().and_then(|n| n.clone());
744 let dpop_proof = dpop_key.generate_proof_with_ath(
745 "POST",
746 &self.token_url,
747 None,
748 nonce.as_deref(),
749 )?;
750 let res = self
751 .http_client
752 .post(&self.token_url)
753 .header("DPoP", dpop_proof)
754 .form(params)
755 .send()
756 .await?;
757 let new_nonce = crate::crypto::dpop_nonce(res.headers());
758 if let Some(n) = &new_nonce {
759 if let Ok(mut slot) = self.as_nonce.lock() {
760 *slot = Some(n.clone());
761 }
762 }
763 if res.status().is_success() {
764 break res;
765 }
766
767 let status = res.status();
768 let body = res.text().await.unwrap_or_default();
769 let wants_nonce = status == reqwest::StatusCode::BAD_REQUEST
770 && serde_json::from_str::<serde_json::Value>(&body)
771 .ok()
772 .and_then(|v| v.get("error").and_then(|e| e.as_str()).map(str::to_string))
773 .as_deref()
774 == Some("use_dpop_nonce");
775 if wants_nonce && new_nonce.is_some() && !retried {
776 info!("Authorization server requires a DPoP nonce; retrying.");
777 retried = true;
778 continue;
779 }
780 error!("Token endpoint returned {}: {}", status, body);
781 return Ok(Err(format!("{status}: {body}")));
782 };
783
784 let tokens: TokenResponse = res.json().await.context("Invalid token response")?;
785 match tokens.token_type.as_deref() {
786 Some(t) if t.eq_ignore_ascii_case("DPoP") => {}
787 other => warn!(
788 "Token endpoint returned token_type {:?} instead of \"DPoP\": the token is not \
789 bound to the DPoP key (RFC 9449 §5).",
790 other
791 ),
792 }
793 Ok(Ok(tokens))
794 }
795
796 fn store_tokens(
799 &self,
800 user_id: &str,
801 tokens: &TokenResponse,
802 scopes: Option<&[String]>,
803 ) -> Result<()> {
804 self.vault.store_token(user_id, &tokens.access_token)?;
805 if let Some(refresh) = &tokens.refresh_token {
807 self.vault.store_refresh_token(user_id, refresh)?;
808 }
809 let scopes = match scopes {
810 Some(s) => s.to_vec(),
811 None => self
812 .vault
813 .get_meta(user_id)?
814 .map(|m| m.scopes)
815 .unwrap_or_default(),
816 };
817 let expires_at = tokens.expires_in.map(|secs| unix_now() + secs);
818 self.vault.store_meta(
819 user_id,
820 &CredentialMeta {
821 issuer: self.issuer.clone(),
822 scopes,
823 expires_at,
824 },
825 )
826 }
827
828 fn bound_meta(&self, user_id: &str) -> Result<Option<CredentialMeta>> {
831 Ok(self
832 .vault
833 .get_meta(user_id)?
834 .filter(|meta| match (&meta.issuer, &self.issuer) {
835 (Some(stored), Some(current)) => stored == current,
836 _ => true,
837 }))
838 }
839
840 pub fn enforce_issuer_binding(&self, user_id: &str) -> Result<()> {
843 let Some(meta) = self.vault.get_meta(user_id)? else {
844 return Ok(());
845 };
846 if self.bound_meta(user_id)?.is_none() {
847 warn!(
848 "Stored credentials were issued by '{}', but the server now uses '{}'; \
849 discarding them.",
850 meta.issuer.as_deref().unwrap_or_default(),
851 self.issuer.as_deref().unwrap_or_default()
852 );
853 self.vault.clear(user_id)?;
854 }
855 Ok(())
856 }
857
858 fn select_scopes(&self, challenged: Option<Vec<String>>, previous: &[String]) -> Vec<String> {
864 let mut scopes: Vec<String> = previous.to_vec();
865 let wanted = challenged
866 .or_else(|| self.resource_scopes.clone())
867 .unwrap_or_default();
868 let offered = |s: &str| {
869 self.as_scopes
870 .as_ref()
871 .is_some_and(|supported| supported.iter().any(|x| x == s))
872 };
873 let mut extra = Vec::new();
874 if self.as_scopes.is_none() || offered("openid") {
876 extra.push("openid".to_string());
877 }
878 if self.request_offline_access && offered("offline_access") {
879 extra.push("offline_access".to_string());
880 }
881 for s in wanted.into_iter().chain(extra) {
882 if !scopes.contains(&s) {
883 scopes.push(s);
884 }
885 }
886 scopes
887 }
888
889 pub fn get_token(&self, user_id: &str) -> Result<Option<String>> {
891 self.vault.get_token(user_id)
892 }
893}
894
895struct LoopbackServer {
897 handle: tokio::task::JoinHandle<()>,
898 shutdown: Option<oneshot::Sender<()>>,
899}
900
901impl LoopbackServer {
902 async fn stop(mut self) {
904 if let Some(tx) = self.shutdown.take() {
905 let _ = tx.send(());
906 }
907 if tokio::time::timeout(Duration::from_secs(2), &mut self.handle)
908 .await
909 .is_err()
910 {
911 self.handle.abort();
912 }
913 }
914}
915
916impl Drop for LoopbackServer {
917 fn drop(&mut self) {
918 self.handle.abort();
919 }
920}
921
922type CallbackResult = std::result::Result<String, String>;
924
925#[derive(Clone)]
926struct AuthServerState {
927 expected_state: String,
928 tx: Arc<tokio::sync::Mutex<Option<oneshot::Sender<CallbackResult>>>>,
929 success_html: Arc<String>,
930 failure_html: Arc<String>,
931 issuer_name: String,
932 resource_name: String,
933 expected_issuer: Option<String>,
935 iss_required: bool,
937}
938
939fn evaluate_callback(query: &AuthCallback, state: &AuthServerState) -> CallbackResult {
941 match (&query.iss, &state.expected_issuer) {
944 (Some(iss), Some(expected)) if iss != expected => {
945 return Err(format!(
946 "issuer mismatch in authorization response ('{iss}', expected '{expected}')"
947 ));
948 }
949 (None, Some(_)) if state.iss_required => {
950 return Err("authorization response is missing the 'iss' parameter".into());
951 }
952 _ => {}
953 }
954 if let Some(error) = &query.error {
955 return Err(match &query.error_description {
956 Some(desc) => format!("{error}: {desc}"),
957 None => error.clone(),
958 });
959 }
960 query
961 .code
962 .clone()
963 .ok_or_else(|| "authorization response has no code".to_string())
964}
965
966async fn handle_callback(
967 query: Query<AuthCallback>,
968 State(state): State<AuthServerState>,
969) -> impl IntoResponse {
970 let render_failure = |status: axum::http::StatusCode, message: &str| {
971 let html = render_template(
972 &state.failure_html,
973 Some(message),
974 &state.issuer_name,
975 &state.resource_name,
976 );
977 (status, axum::response::Html(html)).into_response()
978 };
979
980 if query.state.as_deref() != Some(state.expected_state.as_str()) {
983 return render_failure(axum::http::StatusCode::BAD_REQUEST, "Invalid state");
984 }
985
986 let Some(sender) = state.tx.lock().await.take() else {
987 return render_failure(
988 axum::http::StatusCode::GONE,
989 "Already authenticated or timed out.",
990 );
991 };
992
993 let result = evaluate_callback(&query, &state);
994 let outcome = match &result {
995 Ok(_) => None,
996 Err(reason) => Some(format!("Authorization failed: {reason}")),
997 };
998 let _ = sender.send(result);
999
1000 match outcome {
1001 None => {
1002 let html = render_template(
1003 &state.success_html,
1004 None,
1005 &state.issuer_name,
1006 &state.resource_name,
1007 );
1008 (axum::http::StatusCode::OK, axum::response::Html(html)).into_response()
1009 }
1010 Some(message) => render_failure(axum::http::StatusCode::BAD_REQUEST, &message),
1011 }
1012}
1013
1014fn escape_html(s: &str) -> String {
1015 html_escape::encode_safe(s).to_string()
1016}
1017
1018fn render_template(
1019 template: &str,
1020 error_message: Option<&str>,
1021 issuer_name: &str,
1022 resource_name: &str,
1023) -> String {
1024 let mut result = template.replace("{{ISSUER_NAME}}", &escape_html(issuer_name));
1025 result = result.replace("{{RESOURCE_NAME}}", &escape_html(resource_name));
1026 if let Some(msg) = error_message {
1027 result = result.replace("{{ERROR_MESSAGE}}", &escape_html(msg));
1028 }
1029 result
1030}
1031
1032fn build_authorize_url(auth_endpoint: &str, client_id: &str, request_uri: &str) -> Result<String> {
1035 let mut url = url::Url::parse(auth_endpoint)
1036 .with_context(|| format!("Invalid authorization endpoint '{auth_endpoint}'"))?;
1037 url.query_pairs_mut()
1038 .append_pair("client_id", client_id)
1039 .append_pair("response_type", "code")
1040 .append_pair("request_uri", request_uri);
1041 Ok(url.into())
1042}
1043
1044fn unix_now() -> u64 {
1045 std::time::SystemTime::now()
1046 .duration_since(std::time::UNIX_EPOCH)
1047 .map(|d| d.as_secs())
1048 .unwrap_or_default()
1049}
1050
1051fn is_url_client_id(client_id: &str) -> bool {
1053 client_id.starts_with("https://") || client_id.starts_with("http://")
1054}
1055
1056fn check_cimd_client_id(client_id: &str) -> Result<()> {
1058 let url = url::Url::parse(client_id)
1059 .with_context(|| format!("Invalid client ID metadata document URL '{client_id}'"))?;
1060 if url.scheme() != "https" || url.path().trim_matches('/').is_empty() {
1061 anyhow::bail!(
1062 "Client ID '{}' must be an https URL with a path to be used as a Client ID \
1063 Metadata Document",
1064 client_id
1065 );
1066 }
1067 Ok(())
1068}
1069
1070fn random_urlsafe(len: usize) -> String {
1072 let mut buf = vec![0u8; len];
1073 OsRng.fill_bytes(&mut buf);
1074 URL_SAFE_NO_PAD.encode(buf)
1075}
1076
1077fn pkce_s256_challenge(verifier: &str) -> String {
1079 URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes()))
1080}
1081
1082#[cfg(test)]
1083mod tests {
1084 use super::*;
1085
1086 #[test]
1087 fn test_build_authorize_url_encodes_and_keeps_query() {
1088 let url = build_authorize_url(
1089 "https://as.example.com/authorize?tenant=a",
1090 "my client&x=1",
1091 "urn:ietf:params:oauth:request_uri:abc",
1092 )
1093 .unwrap();
1094 assert_eq!(
1095 url,
1096 "https://as.example.com/authorize?tenant=a&client_id=my+client%26x%3D1\
1097 &response_type=code&request_uri=urn%3Aietf%3Aparams%3Aoauth%3Arequest_uri%3Aabc"
1098 );
1099 assert!(build_authorize_url("not a url", "c", "r").is_err());
1100 }
1101
1102 #[test]
1103 fn test_pkce_s256_challenge() {
1104 assert_eq!(
1106 pkce_s256_challenge("dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"),
1107 "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
1108 );
1109 }
1110
1111 #[test]
1112 fn test_random_urlsafe() {
1113 let a = random_urlsafe(32);
1114 assert_eq!(a.len(), 43); assert!(a
1116 .chars()
1117 .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'));
1118 assert_ne!(a, random_urlsafe(32));
1119 }
1120
1121 #[test]
1122 fn test_escape_html() {
1123 assert_eq!(escape_html("<script>"), "<script>");
1124 assert_eq!(escape_html("a & b"), "a & b");
1125 assert_eq!(
1126 escape_html("\"double quotes\""),
1127 ""double quotes""
1128 );
1129 assert_eq!(escape_html("'single quotes'"), "'single quotes'");
1130 assert_eq!(
1131 escape_html("<img src=x onerror=alert(1)>"),
1132 "<img src=x onerror=alert(1)>"
1133 );
1134 }
1135
1136 #[tokio::test]
1137 async fn test_auth_manager_set_internal_callback_tx() {
1138 let am = AuthManager {
1139 client_id: "c".into(),
1140 auth_url: "a".into(),
1141 token_url: "t".into(),
1142 par_url: "p".into(),
1143 redirect_url: "r".into(),
1144 resource: "res".into(),
1145 http_client: reqwest::Client::new(),
1146 vault: Vault::in_memory("svc_test_set_internal_callback_tx"),
1147 internal_url_tx: Arc::new(tokio::sync::Mutex::new(None)),
1148 internal_callback_tx: Arc::new(tokio::sync::Mutex::new(None)),
1149 issuer_name: "Mock Issuer".into(),
1150 resource_name: "Mock Resource".into(),
1151 success_html: std::sync::Arc::new(crate::templates::DEFAULT_SUCCESS_HTML.to_string()),
1152 failure_html: std::sync::Arc::new(crate::templates::DEFAULT_FAILURE_HTML.to_string()),
1153 timeouts: Timeouts::fast(),
1154 issuer: None,
1155 iss_required: false,
1156 as_nonce: Arc::default(),
1157 as_scopes: None,
1158 resource_scopes: None,
1159 request_offline_access: false,
1160 };
1161
1162 let (tx, _rx) = oneshot::channel::<SocketAddr>();
1163 am.set_internal_callback_tx(tx).await;
1164
1165 let lock = am.internal_callback_tx.lock().await;
1166 assert!(lock.is_some());
1167 }
1168
1169 #[tokio::test]
1170 async fn test_handle_callback_success() {
1171 let (tx, mut rx) = oneshot::channel::<CallbackResult>();
1172 let state = AuthServerState {
1173 expected_state: "test_state".to_string(),
1174 tx: Arc::new(tokio::sync::Mutex::new(Some(tx))),
1175 success_html: std::sync::Arc::new(crate::templates::DEFAULT_SUCCESS_HTML.to_string()),
1176 failure_html: std::sync::Arc::new(crate::templates::DEFAULT_FAILURE_HTML.to_string()),
1177 issuer_name: "Test Issuer".to_string(),
1178 resource_name: "Test Resource".to_string(),
1179 expected_issuer: None,
1180 iss_required: false,
1181 };
1182
1183 let query = Query(AuthCallback {
1184 code: Some("test_code".into()),
1185 state: Some("test_state".into()),
1186 ..Default::default()
1187 });
1188
1189 let response = handle_callback(query, State(state)).await.into_response();
1190 assert_eq!(response.status(), axum::http::StatusCode::OK);
1191 assert_eq!(rx.try_recv().unwrap(), Ok("test_code".to_string()));
1192 }
1193
1194 #[tokio::test]
1195 async fn test_handle_callback_with_templates() -> Result<()> {
1196 let temp_dir = std::env::temp_dir().join(format!("mcp_test_{}", uuid::Uuid::new_v4()));
1197 tokio::fs::create_dir_all(&temp_dir).await?;
1198 tokio::fs::write(temp_dir.join("success.html"), "SUCCESS {{RESOURCE_NAME}}").await?;
1199 tokio::fs::write(temp_dir.join("failure.html"), "FAILURE {{ERROR_MESSAGE}}").await?;
1200
1201 let (tx, mut rx) = oneshot::channel::<CallbackResult>();
1202 let state = AuthServerState {
1203 expected_state: "test_state".to_string(),
1204 tx: Arc::new(tokio::sync::Mutex::new(Some(tx))),
1205 success_html: std::sync::Arc::new(
1206 tokio::fs::read_to_string(temp_dir.join("success.html"))
1207 .await
1208 .unwrap(),
1209 ),
1210 failure_html: std::sync::Arc::new(
1211 tokio::fs::read_to_string(temp_dir.join("failure.html"))
1212 .await
1213 .unwrap(),
1214 ),
1215 issuer_name: "Test Issuer".to_string(),
1216 resource_name: "Test Resource".to_string(),
1217 expected_issuer: None,
1218 iss_required: false,
1219 };
1220
1221 let query_ok = Query(AuthCallback {
1223 code: Some("test_code".into()),
1224 state: Some("test_state".into()),
1225 ..Default::default()
1226 });
1227 let res_ok = handle_callback(query_ok, State(state.clone()))
1228 .await
1229 .into_response();
1230 assert_eq!(res_ok.status(), axum::http::StatusCode::OK);
1231 let body_ok = axum::body::to_bytes(res_ok.into_body(), 1024)
1232 .await
1233 .unwrap();
1234 assert!(String::from_utf8_lossy(&body_ok).contains("SUCCESS Test Resource"));
1235 assert_eq!(rx.try_recv().unwrap(), Ok("test_code".to_string()));
1236
1237 let query_err = Query(AuthCallback {
1239 code: Some("c".into()),
1240 state: Some("wrong".into()),
1241 ..Default::default()
1242 });
1243 let res_err = handle_callback(query_err, State(state.clone()))
1244 .await
1245 .into_response();
1246 assert_eq!(res_err.status(), axum::http::StatusCode::BAD_REQUEST);
1247 let body_err = axum::body::to_bytes(res_err.into_body(), 1024)
1248 .await
1249 .unwrap();
1250 assert!(String::from_utf8_lossy(&body_err).contains("FAILURE Invalid state"));
1251
1252 let query_gone = Query(AuthCallback {
1254 code: Some("c".into()),
1255 state: Some("test_state".into()),
1256 ..Default::default()
1257 });
1258 let res_gone = handle_callback(query_gone, State(state.clone()))
1259 .await
1260 .into_response();
1261 assert_eq!(res_gone.status(), axum::http::StatusCode::GONE);
1262 let body_gone = axum::body::to_bytes(res_gone.into_body(), 1024)
1263 .await
1264 .unwrap();
1265 assert!(String::from_utf8_lossy(&body_gone).contains("FAILURE Already authenticated"));
1266
1267 tokio::fs::remove_dir_all(temp_dir).await?;
1268 Ok(())
1269 }
1270
1271 #[tokio::test]
1272 async fn test_handle_callback_invalid_state() {
1273 let (tx, _rx) = oneshot::channel::<CallbackResult>();
1274 let state = AuthServerState {
1275 expected_state: "expected".to_string(),
1276 tx: Arc::new(tokio::sync::Mutex::new(Some(tx))),
1277 success_html: std::sync::Arc::new(crate::templates::DEFAULT_SUCCESS_HTML.to_string()),
1278 failure_html: std::sync::Arc::new(crate::templates::DEFAULT_FAILURE_HTML.to_string()),
1279 issuer_name: "Test Issuer".to_string(),
1280 resource_name: "Test Resource".to_string(),
1281 expected_issuer: None,
1282 iss_required: false,
1283 };
1284
1285 let query = Query(AuthCallback {
1286 code: Some("code".into()),
1287 state: Some("wrong".into()),
1288 ..Default::default()
1289 });
1290
1291 let response = handle_callback(query, State(state)).await.into_response();
1292 assert_eq!(response.status(), axum::http::StatusCode::BAD_REQUEST);
1293 }
1294
1295 #[tokio::test]
1296 async fn test_handle_callback_xss_prevention() -> Result<()> {
1297 let (tx, _rx) = oneshot::channel::<CallbackResult>();
1298 let state = AuthServerState {
1299 expected_state: "test_state".to_string(),
1300 tx: Arc::new(tokio::sync::Mutex::new(Some(tx))),
1301 success_html: std::sync::Arc::new(crate::templates::DEFAULT_SUCCESS_HTML.to_string()),
1302 failure_html: std::sync::Arc::new(crate::templates::DEFAULT_FAILURE_HTML.to_string()),
1303 issuer_name: "<script>alert('xss')</script>".to_string(),
1304 resource_name: "<b>Bold Resource</b>".to_string(),
1305 expected_issuer: None,
1306 iss_required: false,
1307 };
1308
1309 let query_err = Query(AuthCallback {
1311 code: Some("c".into()),
1312 state: Some("wrong".into()),
1313 ..Default::default()
1314 });
1315 let res_err = handle_callback(query_err, State(state.clone()))
1316 .await
1317 .into_response();
1318 let body_err = axum::body::to_bytes(res_err.into_body(), 4096)
1319 .await
1320 .unwrap();
1321 let html_err = String::from_utf8_lossy(&body_err);
1322
1323 assert!(html_err.contains("<script>alert('xss')</script>"));
1324 assert!(html_err.contains("<b>Bold Resource</b>"));
1325 assert!(!html_err.contains("<script>"));
1326 assert!(!html_err.contains("<b>"));
1327
1328 Ok(())
1329 }
1330
1331 fn callback_state(
1332 tx: oneshot::Sender<CallbackResult>,
1333 expected_issuer: Option<&str>,
1334 iss_required: bool,
1335 ) -> AuthServerState {
1336 AuthServerState {
1337 expected_state: "st".to_string(),
1338 tx: Arc::new(tokio::sync::Mutex::new(Some(tx))),
1339 success_html: Arc::new("OK".to_string()),
1340 failure_html: Arc::new("FAIL {{ERROR_MESSAGE}}".to_string()),
1341 issuer_name: "Issuer".to_string(),
1342 resource_name: "Resource".to_string(),
1343 expected_issuer: expected_issuer.map(str::to_string),
1344 iss_required,
1345 }
1346 }
1347
1348 async fn body_of(res: axum::response::Response) -> String {
1349 let bytes = axum::body::to_bytes(res.into_body(), 4096).await.unwrap();
1350 String::from_utf8_lossy(&bytes).to_string()
1351 }
1352
1353 #[tokio::test]
1354 async fn test_handle_callback_error_fails_flow_immediately() {
1355 let (tx, mut rx) = oneshot::channel::<CallbackResult>();
1356 let query = Query(AuthCallback {
1357 state: Some("st".into()),
1358 error: Some("access_denied".into()),
1359 error_description: Some("User <denied>".into()),
1360 ..Default::default()
1361 });
1362 let res = handle_callback(query, State(callback_state(tx, None, false)))
1363 .await
1364 .into_response();
1365 assert_eq!(res.status(), axum::http::StatusCode::BAD_REQUEST);
1366 assert!(body_of(res)
1367 .await
1368 .contains("access_denied: User <denied>"));
1369 assert_eq!(
1370 rx.try_recv().unwrap(),
1371 Err("access_denied: User <denied>".to_string())
1372 );
1373 }
1374
1375 #[tokio::test]
1376 async fn test_handle_callback_error_with_wrong_state_is_ignored() {
1377 let (tx, mut rx) = oneshot::channel::<CallbackResult>();
1378 let query = Query(AuthCallback {
1379 state: Some("other".into()),
1380 error: Some("access_denied".into()),
1381 ..Default::default()
1382 });
1383 let res = handle_callback(query, State(callback_state(tx, None, false)))
1384 .await
1385 .into_response();
1386 assert_eq!(res.status(), axum::http::StatusCode::BAD_REQUEST);
1387 assert!(rx.try_recv().is_err());
1389 }
1390
1391 #[tokio::test]
1392 async fn test_handle_callback_missing_state_is_ignored() {
1393 let (tx, mut rx) = oneshot::channel::<CallbackResult>();
1394 let query = Query(AuthCallback {
1395 code: Some("c".into()),
1396 ..Default::default()
1397 });
1398 let res = handle_callback(query, State(callback_state(tx, None, false)))
1399 .await
1400 .into_response();
1401 assert_eq!(res.status(), axum::http::StatusCode::BAD_REQUEST);
1402 assert!(rx.try_recv().is_err());
1403 }
1404
1405 #[test]
1406 fn test_evaluate_callback_issuer_checks() {
1407 let ok = |iss: Option<&str>, expected: Option<&str>, required: bool| {
1408 let (tx, _rx) = oneshot::channel::<CallbackResult>();
1409 let query = AuthCallback {
1410 code: Some("c".into()),
1411 state: Some("st".into()),
1412 iss: iss.map(str::to_string),
1413 ..Default::default()
1414 };
1415 evaluate_callback(&query, &callback_state(tx, expected, required))
1416 };
1417 let issuer = Some("https://as.example.com");
1418 assert_eq!(ok(issuer, issuer, true), Ok("c".to_string()));
1419 assert_eq!(ok(issuer, issuer, false), Ok("c".to_string()));
1420 assert_eq!(ok(None, issuer, false), Ok("c".to_string()));
1421 assert_eq!(ok(None, None, true), Ok("c".to_string()));
1422 assert!(ok(Some("https://evil.example.com"), issuer, false)
1423 .unwrap_err()
1424 .contains("issuer mismatch"));
1425 assert!(ok(None, issuer, true)
1426 .unwrap_err()
1427 .contains("missing the 'iss'"));
1428 }
1429
1430 #[test]
1431 fn test_evaluate_callback_issuer_mismatch_beats_error() {
1432 let (tx, _rx) = oneshot::channel::<CallbackResult>();
1433 let query = AuthCallback {
1434 state: Some("st".into()),
1435 error: Some("access_denied".into()),
1436 iss: Some("https://evil.example.com".into()),
1437 ..Default::default()
1438 };
1439 let state = callback_state(tx, Some("https://as.example.com"), false);
1440 assert!(evaluate_callback(&query, &state)
1441 .unwrap_err()
1442 .contains("issuer mismatch"));
1443 }
1444
1445 #[test]
1446 fn test_evaluate_callback_requires_code() {
1447 let (tx, _rx) = oneshot::channel::<CallbackResult>();
1448 let query = AuthCallback {
1449 state: Some("st".into()),
1450 ..Default::default()
1451 };
1452 assert!(evaluate_callback(&query, &callback_state(tx, None, false)).is_err());
1453 }
1454
1455 #[tokio::test]
1456 async fn test_loopback_servers_on_the_same_port_dont_share_connections() -> Result<()> {
1457 let port = tokio::net::TcpListener::bind("127.0.0.1:0")
1460 .await?
1461 .local_addr()?
1462 .port();
1463 let am = AuthManager {
1464 client_id: "c".into(),
1465 auth_url: "http://localhost/auth".into(),
1466 token_url: "http://localhost/token".into(),
1467 par_url: "http://localhost/par".into(),
1468 redirect_url: format!("http://127.0.0.1:{port}/callback"),
1469 resource: "res".into(),
1470 http_client: reqwest::Client::new(),
1471 vault: Vault::in_memory("svc"),
1472 internal_url_tx: Arc::new(tokio::sync::Mutex::new(None)),
1473 internal_callback_tx: Arc::new(tokio::sync::Mutex::new(None)),
1474 issuer_name: "Mock Issuer".into(),
1475 resource_name: "Mock Resource".into(),
1476 success_html: Arc::new("OK".into()),
1477 failure_html: Arc::new("FAIL".into()),
1478 timeouts: Timeouts::fast(),
1479 issuer: None,
1480 iss_required: false,
1481 as_nonce: Arc::default(),
1482 as_scopes: None,
1483 resource_scopes: None,
1484 request_offline_access: false,
1485 };
1486 let browser = reqwest::Client::new(); let callback = |state: &str| {
1488 format!("http://127.0.0.1:{port}/callback?state={state}&code=code-{state}")
1489 };
1490
1491 for state in ["first", "second"] {
1492 let (server, rx) = am.setup_loopback_server(state.to_string()).await?;
1493 let resp = browser.get(callback(state)).send().await?;
1494 assert_eq!(
1495 resp.status(),
1496 axum::http::StatusCode::OK,
1497 "callback for {state}"
1498 );
1499 assert_eq!(resp.headers()["connection"], "close");
1500 assert_eq!(rx.await?, Ok(format!("code-{state}")));
1501 server.stop().await;
1502 }
1503 Ok(())
1504 }
1505
1506 #[tokio::test]
1507 async fn test_reauthenticate_rejects_non_loopback_redirect() {
1508 let am = AuthManager {
1509 client_id: "c".into(),
1510 auth_url: "http://localhost/auth".into(),
1511 token_url: "http://localhost/token".into(),
1512 par_url: "http://localhost/par".into(),
1513 redirect_url: "http://0.0.0.0:8082/callback".into(),
1514 resource: "res".into(),
1515 http_client: reqwest::Client::new(),
1516 vault: Vault::in_memory("svc"),
1517 internal_url_tx: Arc::new(tokio::sync::Mutex::new(None)),
1518 internal_callback_tx: Arc::new(tokio::sync::Mutex::new(None)),
1519 issuer_name: "Mock Issuer".into(),
1520 resource_name: "Mock Resource".into(),
1521 success_html: Arc::new(String::new()),
1522 failure_html: Arc::new(String::new()),
1523 timeouts: Timeouts::fast(),
1524 issuer: None,
1525 iss_required: false,
1526 as_nonce: Arc::default(),
1527 as_scopes: None,
1528 resource_scopes: None,
1529 request_offline_access: false,
1530 };
1531 let err = am.reauthenticate("user", None, None).await.unwrap_err();
1532 assert!(err.to_string().contains("RFC 8252"), "{err}");
1533 }
1534
1535 #[tokio::test]
1536 async fn test_failed_flow_keeps_existing_credentials() {
1537 let vault = Vault::in_memory("svc");
1539 vault.store_token("user", "old-token").unwrap();
1540 vault.store_dpop_key("user", &[7u8; 32]).unwrap();
1541 let am = AuthManager {
1542 client_id: "c".into(),
1543 auth_url: "http://localhost:1/auth".into(),
1544 token_url: "http://localhost:1/token".into(),
1545 par_url: "http://localhost:1/par".into(),
1546 redirect_url: "http://127.0.0.1:0/callback".into(),
1547 resource: "res".into(),
1548 http_client: reqwest::Client::new(),
1549 vault: vault.clone(),
1550 internal_url_tx: Arc::new(tokio::sync::Mutex::new(None)),
1551 internal_callback_tx: Arc::new(tokio::sync::Mutex::new(None)),
1552 issuer_name: "Mock Issuer".into(),
1553 resource_name: "Mock Resource".into(),
1554 success_html: Arc::new(String::new()),
1555 failure_html: Arc::new(String::new()),
1556 timeouts: Timeouts::fast(),
1557 issuer: None,
1558 iss_required: false,
1559 as_nonce: Arc::default(),
1560 as_scopes: None,
1561 resource_scopes: None,
1562 request_offline_access: false,
1563 };
1564 assert!(am.reauthenticate("user", None, None).await.is_err());
1565 assert_eq!(vault.get_dpop_key("user").unwrap(), Some(vec![7u8; 32]));
1566 assert_eq!(vault.get_token("user").unwrap(), Some("old-token".into()));
1567 }
1568
1569 #[tokio::test]
1570 async fn test_discover_rejects_insecure_endpoints() {
1571 let config = OidcConfig {
1572 client_id: "c".into(),
1573 redirect_url: "http://127.0.0.1:1/callback".into(),
1574 auth_url_override: Some("http://as.example.com/auth".into()),
1575 token_url_override: Some("https://as.example.com/token".into()),
1576 par_url_override: Some("https://as.example.com/par".into()),
1577 ..Default::default()
1578 };
1579 let err = AuthManager::discover(
1580 config.clone(),
1581 "https://mcp.example.com".into(),
1582 Vault::in_memory("svc"),
1583 None,
1584 )
1585 .await
1586 .err()
1587 .unwrap();
1588 assert!(err.to_string().contains("must use HTTPS"), "{err}");
1589
1590 let allowed = OidcConfig {
1591 allow_insecure_http: true,
1592 ..config
1593 };
1594 let am = AuthManager::discover(
1597 allowed,
1598 "http://127.0.0.1:1/rpc".into(),
1599 Vault::in_memory("svc"),
1600 None,
1601 )
1602 .await
1603 .unwrap();
1604 assert_eq!(am.auth_url, "http://as.example.com/auth");
1605 }
1606
1607 #[tokio::test]
1608 async fn test_auth_manager_set_internal_url_tx() {
1609 let am = AuthManager {
1610 client_id: "c".into(),
1611 auth_url: "a".into(),
1612 token_url: "t".into(),
1613 par_url: "p".into(),
1614 redirect_url: "r".into(),
1615 resource: "res".into(),
1616 http_client: reqwest::Client::new(),
1617 vault: Vault::in_memory("svc_test_set_internal_url_tx"),
1618 internal_url_tx: Arc::new(tokio::sync::Mutex::new(None)),
1619 internal_callback_tx: Arc::new(tokio::sync::Mutex::new(None)),
1620 issuer_name: "Mock Issuer".into(),
1621 resource_name: "Mock Resource".into(),
1622 success_html: std::sync::Arc::new(crate::templates::DEFAULT_SUCCESS_HTML.to_string()),
1623 failure_html: std::sync::Arc::new(crate::templates::DEFAULT_FAILURE_HTML.to_string()),
1624 timeouts: Timeouts::fast(),
1625 issuer: None,
1626 iss_required: false,
1627 as_nonce: Arc::default(),
1628 as_scopes: None,
1629 resource_scopes: None,
1630 request_offline_access: false,
1631 };
1632
1633 let (tx, _rx) = oneshot::channel::<String>();
1634 am.set_internal_url_tx(tx).await;
1635
1636 let lock = am.internal_url_tx.lock().await;
1637 assert!(lock.is_some());
1638 }
1639
1640 #[tokio::test]
1641 async fn test_auth_manager_get_token_fresh() -> Result<()> {
1642 let am = AuthManager {
1643 client_id: "c".into(),
1644 auth_url: "a".into(),
1645 token_url: "t".into(),
1646 par_url: "p".into(),
1647 redirect_url: "r".into(),
1648 resource: "res".into(),
1649 http_client: reqwest::Client::new(),
1650 vault: Vault::in_memory("svc_test_get_token"),
1651 internal_url_tx: Arc::new(tokio::sync::Mutex::new(None)),
1652 internal_callback_tx: Arc::new(tokio::sync::Mutex::new(None)),
1653 issuer_name: "Mock Issuer".into(),
1654 resource_name: "Mock Resource".into(),
1655 success_html: std::sync::Arc::new(crate::templates::DEFAULT_SUCCESS_HTML.to_string()),
1656 failure_html: std::sync::Arc::new(crate::templates::DEFAULT_FAILURE_HTML.to_string()),
1657 timeouts: Timeouts::fast(),
1658 issuer: None,
1659 iss_required: false,
1660 as_nonce: Arc::default(),
1661 as_scopes: None,
1662 resource_scopes: None,
1663 request_offline_access: false,
1664 };
1665 am.vault.store_token("user", "token")?;
1666
1667 assert_eq!(am.get_token("user")?, Some("token".into()));
1668 Ok(())
1669 }
1670
1671 #[tokio::test]
1672 async fn test_auth_manager_discover_failure() {
1673 let config = OidcConfig {
1674 discovery_url: Some("http://localhost:1/invalid".into()),
1675 client_id: "c".into(),
1676 redirect_url: "r".into(),
1677 timeouts: crate::auth::Timeouts::fast(),
1678 ..Default::default()
1679 };
1680 let res =
1681 AuthManager::discover(config, "res".to_string(), Vault::in_memory("svc"), None).await;
1682 assert!(res.is_err());
1683 }
1684
1685 #[tokio::test]
1686 async fn test_discover_prefers_configured_discovery_url() -> Result<()> {
1687 let app = Router::new().route(
1688 "/oidc",
1689 get(|| async {
1690 axum::Json(serde_json::json!({
1691 "issuer": "https://configured.example.com",
1692 "authorization_endpoint": "https://configured.example.com/auth",
1693 "token_endpoint": "https://configured.example.com/token",
1694 "pushed_authorization_request_endpoint": "https://configured.example.com/par",
1695 "code_challenge_methods_supported": ["S256"]
1696 }))
1697 }),
1698 );
1699 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
1700 let base = format!("http://{}", listener.local_addr()?);
1701 tokio::spawn(async move {
1702 let _ = axum::serve(listener, app).await;
1703 });
1704
1705 let config = OidcConfig {
1706 discovery_url: Some(format!("{base}/oidc")),
1707 client_id: "c".into(),
1708 redirect_url: "http://127.0.0.1:1/callback".into(),
1709 ..Default::default()
1710 };
1711 let am = AuthManager::discover(
1714 config,
1715 format!("{base}/rpc"),
1716 Vault::in_memory("svc"),
1717 Some(&format!("{base}/missing")),
1718 )
1719 .await?;
1720 assert_eq!(am.token_url, "https://configured.example.com/token");
1721 assert_eq!(am.par_url, "https://configured.example.com/par");
1722 Ok(())
1723 }
1724
1725 async fn discover_with(doc: serde_json::Value, config: OidcConfig) -> Result<AuthManager> {
1727 let app = Router::new().route(
1728 "/oidc",
1729 get(move || {
1730 let doc = doc.clone();
1731 async move { axum::Json(doc) }
1732 }),
1733 );
1734 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
1735 let base = format!("http://{}", listener.local_addr()?);
1736 tokio::spawn(async move {
1737 let _ = axum::serve(listener, app).await;
1738 });
1739 AuthManager::discover(
1740 OidcConfig {
1741 discovery_url: Some(format!("{base}/oidc")),
1742 client_id: "c".into(),
1743 redirect_url: "http://127.0.0.1:1/callback".into(),
1744 ..config
1745 },
1746 format!("{base}/rpc"),
1747 Vault::in_memory("svc"),
1748 None,
1749 )
1750 .await
1751 }
1752
1753 fn as_metadata(extra: serde_json::Value) -> serde_json::Value {
1754 let mut doc = serde_json::json!({
1755 "issuer": "https://as.example.com",
1756 "authorization_endpoint": "https://as.example.com/auth",
1757 "token_endpoint": "https://as.example.com/token",
1758 "pushed_authorization_request_endpoint": "https://as.example.com/par",
1759 "code_challenge_methods_supported": ["S256"]
1760 });
1761 for (k, v) in extra.as_object().unwrap() {
1762 doc[k] = v.clone();
1763 }
1764 doc
1765 }
1766
1767 #[tokio::test]
1768 async fn test_discover_requires_pkce_s256() {
1769 for methods in [serde_json::Value::Null, serde_json::json!(["plain"])] {
1770 let doc = as_metadata(serde_json::json!({"code_challenge_methods_supported": methods}));
1771 let err = discover_with(doc, OidcConfig::default())
1772 .await
1773 .err()
1774 .unwrap();
1775 assert!(err.to_string().contains("PKCE S256"), "{err}");
1776 }
1777 }
1778
1779 #[tokio::test]
1780 async fn test_discover_checks_expected_issuer() -> Result<()> {
1781 let pinned = |issuer: &str| OidcConfig {
1782 expected_issuer: Some(issuer.into()),
1783 ..Default::default()
1784 };
1785 let err = discover_with(
1786 as_metadata(serde_json::json!({})),
1787 pinned("https://other.example.com"),
1788 )
1789 .await
1790 .err()
1791 .unwrap();
1792 assert!(err.to_string().contains("--oidc-issuer"), "{err}");
1793 let am = discover_with(
1794 as_metadata(serde_json::json!({})),
1795 pinned("https://as.example.com"),
1796 )
1797 .await?;
1798 assert_eq!(am.issuer.as_deref(), Some("https://as.example.com"));
1799 Ok(())
1800 }
1801
1802 #[tokio::test]
1803 async fn test_discover_rejects_invalid_cimd_client_id() {
1804 let res = AuthManager::discover(
1805 OidcConfig {
1806 client_id: "http://client.example.com/meta.json".into(),
1807 redirect_url: "http://127.0.0.1:1/callback".into(),
1808 auth_url_override: Some("https://as/auth".into()),
1809 token_url_override: Some("https://as/token".into()),
1810 par_url_override: Some("https://as/par".into()),
1811 ..Default::default()
1812 },
1813 "http://127.0.0.1:1/rpc".into(),
1814 Vault::in_memory("svc"),
1815 None,
1816 )
1817 .await;
1818 assert!(res
1819 .err()
1820 .unwrap()
1821 .to_string()
1822 .contains("https URL with a path"));
1823 assert!(check_cimd_client_id("https://client.example.com/oauth/meta.json").is_ok());
1824 assert!(check_cimd_client_id("https://client.example.com/").is_err());
1825 }
1826
1827 #[tokio::test]
1828 async fn test_select_scopes() -> Result<()> {
1829 let mut am = discover_with(
1830 as_metadata(
1831 serde_json::json!({"scopes_supported": ["openid", "offline_access", "mcp:read"]}),
1832 ),
1833 OidcConfig {
1834 request_offline_access: true,
1835 ..Default::default()
1836 },
1837 )
1838 .await?;
1839 let v = |s: &[&str]| s.iter().map(|x| x.to_string()).collect::<Vec<_>>();
1840
1841 assert_eq!(
1843 am.select_scopes(Some(v(&["files:read"])), &[]),
1844 v(&["files:read", "openid", "offline_access"])
1845 );
1846 assert_eq!(
1848 am.select_scopes(Some(v(&["files:write"])), &v(&["files:read", "openid"])),
1849 v(&["files:read", "openid", "files:write", "offline_access"])
1850 );
1851 am.resource_scopes = Some(v(&["mcp:read"]));
1853 assert_eq!(
1854 am.select_scopes(None, &[]),
1855 v(&["mcp:read", "openid", "offline_access"])
1856 );
1857 am.as_scopes = Some(v(&["mcp:read"]));
1859 assert_eq!(am.select_scopes(None, &[]), v(&["mcp:read"]));
1860 am.as_scopes = Some(v(&["openid", "offline_access"]));
1862 am.request_offline_access = false;
1863 assert_eq!(am.select_scopes(None, &[]), v(&["mcp:read", "openid"]));
1864 am.as_scopes = None;
1866 am.resource_scopes = None;
1867 assert_eq!(am.select_scopes(None, &[]), v(&["openid"]));
1868 Ok(())
1869 }
1870
1871 #[tokio::test]
1872 async fn test_issuer_binding() -> Result<()> {
1873 let am = discover_with(as_metadata(serde_json::json!({})), OidcConfig::default()).await?;
1874 let other = CredentialMeta {
1875 issuer: Some("https://previous-as.example.com".into()),
1876 ..Default::default()
1877 };
1878
1879 am.vault.store_token("u", "t")?;
1881 am.vault.store_refresh_token("u", "r")?;
1882 am.vault.store_meta("u", &other)?;
1883 am.enforce_issuer_binding("u")?;
1884 assert_eq!(am.vault.get_token("u")?, None);
1885 assert_eq!(am.vault.get_refresh_token("u")?, None);
1886
1887 am.vault.store_refresh_token("u", "r")?;
1889 am.vault.store_meta("u", &other)?;
1890 assert!(!am.refresh("u").await?);
1891 assert_eq!(am.vault.get_refresh_token("u")?, None);
1892
1893 am.vault.store_token("u", "t")?;
1895 am.vault.store_meta(
1896 "u",
1897 &CredentialMeta {
1898 issuer: Some("https://as.example.com".into()),
1899 ..Default::default()
1900 },
1901 )?;
1902 am.enforce_issuer_binding("u")?;
1903 assert_eq!(am.vault.get_token("u")?, Some("t".into()));
1904 Ok(())
1905 }
1906
1907 #[tokio::test]
1908 async fn test_discover_overrides_skip_network() -> Result<()> {
1909 let config = OidcConfig {
1910 client_id: "c".into(),
1911 redirect_url: "http://127.0.0.1:1/callback".into(),
1912 auth_url_override: Some("https://as/auth".into()),
1913 token_url_override: Some("https://as/token".into()),
1914 par_url_override: Some("https://as/par".into()),
1915 ..Default::default()
1916 };
1917 let am = AuthManager::discover(
1918 config,
1919 "http://127.0.0.1:1/rpc".into(),
1920 Vault::in_memory("svc"),
1921 None,
1922 )
1923 .await?;
1924 assert_eq!(am.issuer_name, "Custom Provider");
1925 assert_eq!(am.resource_name, "http://127.0.0.1:1/rpc");
1926 Ok(())
1927 }
1928
1929 #[tokio::test]
1930 async fn test_auth_manager_manual_token_exchange_failure() -> Result<()> {
1931 let am = AuthManager {
1932 client_id: "c".into(),
1933 auth_url: "a".into(),
1934 token_url: "http://localhost:1/token".into(),
1935 par_url: "p".into(),
1936 redirect_url: "r".into(),
1937 resource: "res".into(),
1938 http_client: reqwest::Client::new(),
1939 vault: Vault::in_memory("svc_test_token_fail"),
1940 internal_url_tx: Arc::new(tokio::sync::Mutex::new(None)),
1941 internal_callback_tx: Arc::new(tokio::sync::Mutex::new(None)),
1942 issuer_name: "Mock Issuer".into(),
1943 resource_name: "Mock Resource".into(),
1944 success_html: std::sync::Arc::new(crate::templates::DEFAULT_SUCCESS_HTML.to_string()),
1945 failure_html: std::sync::Arc::new(crate::templates::DEFAULT_FAILURE_HTML.to_string()),
1946 timeouts: Timeouts::fast(),
1947 issuer: None,
1948 iss_required: false,
1949 as_nonce: Arc::default(),
1950 as_scopes: None,
1951 resource_scopes: None,
1952 request_offline_access: false,
1953 };
1954 let key = crate::crypto::DpopKey::generate();
1955 let res = am
1956 .manual_token_exchange("user", "code", "verifier", &key, &[])
1957 .await;
1958 assert!(res.is_err());
1959 Ok(())
1960 }
1961
1962 #[tokio::test]
1963 async fn test_auth_manager_reauthenticate_addr_in_use() -> Result<()> {
1964 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
1965 let addr = listener.local_addr()?;
1966
1967 let am = AuthManager {
1968 client_id: "c".into(),
1969 auth_url: "http://localhost/auth".into(),
1970 token_url: "http://localhost/token".into(),
1971 par_url: "http://localhost/par".into(),
1972 redirect_url: format!("http://127.0.0.1:{}/callback", addr.port()),
1973 resource: "res".into(),
1974 http_client: reqwest::Client::new(),
1975 vault: Vault::in_memory("svc_test_addr_in_use"),
1976 internal_url_tx: Arc::new(tokio::sync::Mutex::new(None)),
1977 internal_callback_tx: Arc::new(tokio::sync::Mutex::new(None)),
1978 issuer_name: "Mock Issuer".into(),
1979 resource_name: "Mock Resource".into(),
1980 success_html: std::sync::Arc::new(crate::templates::DEFAULT_SUCCESS_HTML.to_string()),
1981 failure_html: std::sync::Arc::new(crate::templates::DEFAULT_FAILURE_HTML.to_string()),
1982 timeouts: Timeouts::fast(),
1983 issuer: None,
1984 iss_required: false,
1985 as_nonce: Arc::default(),
1986 as_scopes: None,
1987 resource_scopes: None,
1988 request_offline_access: false,
1989 };
1990
1991 let res = tokio::time::timeout(
1993 std::time::Duration::from_secs(1),
1994 am.reauthenticate("user", None, None),
1995 )
1996 .await?;
1997 assert!(res.is_err());
1998 assert!(res.err().unwrap().to_string().contains("Failed to bind"));
1999 Ok(())
2000 }
2001
2002 #[tokio::test]
2003 async fn test_auth_manager_reauthenticate_timeout() -> Result<()> {
2004 let par_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
2005 let par_addr = par_listener.local_addr()?;
2006 let par_url = format!("http://127.0.0.1:{}/par", par_addr.port());
2007
2008 let cb_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
2009 let cb_addr = cb_listener.local_addr()?;
2010 drop(cb_listener); let am = AuthManager {
2013 client_id: "c".into(),
2014 auth_url: "http://localhost/auth".into(),
2015 token_url: "http://localhost/token".into(),
2016 par_url: par_url.clone(),
2017 redirect_url: format!("http://127.0.0.1:{}/callback", cb_addr.port()),
2018 resource: "res".into(),
2019 http_client: reqwest::Client::new(),
2020 vault: Vault::in_memory("svc_test_reauth_timeout"),
2021 internal_url_tx: Arc::new(tokio::sync::Mutex::new(None)),
2022 internal_callback_tx: Arc::new(tokio::sync::Mutex::new(None)),
2023 issuer_name: "Mock Issuer".into(),
2024 resource_name: "Mock Resource".into(),
2025 success_html: std::sync::Arc::new(crate::templates::DEFAULT_SUCCESS_HTML.to_string()),
2026 failure_html: std::sync::Arc::new(crate::templates::DEFAULT_FAILURE_HTML.to_string()),
2027 timeouts: Timeouts::fast(),
2028 issuer: None,
2029 iss_required: false,
2030 as_nonce: Arc::default(),
2031 as_scopes: None,
2032 resource_scopes: None,
2033 request_offline_access: false,
2034 };
2035
2036 let par_app = Router::new().route(
2038 "/par",
2039 axum::routing::post(|| async move {
2040 axum::Json(serde_json::json!({
2041 "request_uri": "urn:ietf:params:oauth:request_uri:123",
2042 "expires_in": 3600
2043 }))
2044 }),
2045 );
2046 tokio::spawn(async move {
2047 let _ = axum::serve(par_listener, par_app).await;
2048 });
2049
2050 std::env::set_var("MCP_AIRLOCK_SKIP_OPEN_BROWSER", "1");
2051
2052 let res = am.reauthenticate("user", None, None).await;
2053 assert!(res.is_err());
2054 let err_msg = res.err().unwrap().to_string();
2055 assert!(err_msg.contains("Authentication timed out"));
2056 Ok(())
2057 }
2058}