Skip to main content

mcp_airlock/
auth.rs

1//! # OIDC/OAuth2 Authentication Manager
2//!
3//! This module implements the OpenID Connect (OIDC) flow with FAPI 2.0 security
4//! enhancements, including:
5//! - **Pushed Authorization Requests (PAR)**
6//! - **PKCE** (Proof Key for Code Exchange)
7//! - **DPoP** (Demonstrating Proof-of-Possession)
8//! - **RFC 8707 Resource Indicators**
9//!
10//! It also includes a local loopback server to handle the OAuth2 callback.
11
12use 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/// Configuration for the OIDC provider and local callback server.
37#[derive(Clone)]
38pub struct OidcConfig {
39    /// URL for OIDC discovery (.well-known/openid-configuration).
40    pub discovery_url: Option<String>,
41    /// OIDC Client ID.
42    pub client_id: String,
43    /// URL for the OAuth2 callback (must match provider configuration).
44    pub redirect_url: String,
45    /// Optional override for the authorization endpoint.
46    pub auth_url_override: Option<String>,
47    /// Optional override for the token endpoint.
48    pub token_url_override: Option<String>,
49    /// Optional override for the PAR endpoint.
50    pub par_url_override: Option<String>,
51    /// Internal channel to communicate the auth URL (used for automation/tests).
52    pub internal_url_tx: Arc<tokio::sync::Mutex<Option<oneshot::Sender<String>>>>,
53    /// Internal channel to communicate the callback server address (used for automation/tests).
54    pub internal_callback_tx: Arc<tokio::sync::Mutex<Option<oneshot::Sender<SocketAddr>>>>,
55    /// Directory containing custom templates for success/failure pages.
56    pub template_dir: Option<std::path::PathBuf>,
57    /// Timeouts and retry delays used by the auth flow and the SSE listener.
58    pub timeouts: Timeouts,
59    /// Accept plain-HTTP authorization server endpoints on non-loopback hosts.
60    pub allow_insecure_http: bool,
61    /// The issuer the (pre-registered) client belongs to. Discovery must find
62    /// exactly this authorization server.
63    pub expected_issuer: Option<String>,
64    /// Ask for `offline_access` when the authorization server offers it.
65    /// Off by default: providers list it even when the client may not use it
66    /// (Keycloak then fails the code exchange with `not_allowed`).
67    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/// Timeouts and retry delays. `Default` holds the production values.
91#[derive(Clone, Debug, PartialEq, Eq)]
92pub struct Timeouts {
93    /// How long to wait for the user to complete the browser login.
94    pub auth: Duration,
95    /// Delay between attempts to bind the loopback callback port.
96    pub bind_retry: Duration,
97    /// Base delay before reconnecting a dropped SSE stream.
98    pub sse_retry_base: Duration,
99    /// Maximum random jitter added to `sse_retry_base`.
100    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    /// Short timeouts, intended for tests.
116    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
126/// Connect timeout for every outbound HTTP connection.
127pub(crate) const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
128/// Total timeout for requests to the authorization server (discovery, PAR, token).
129const 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/// Manages OIDC discovery, token exchange, and the DPoP flow.
146#[derive(Clone)]
147pub struct AuthManager {
148    /// OIDC Client ID.
149    client_id: String,
150    /// URL for the authorization endpoint.
151    auth_url: String,
152    /// URL for the token endpoint.
153    token_url: String,
154    /// URL for the PAR endpoint.
155    par_url: String,
156    /// URL for the OAuth2 callback.
157    redirect_url: String,
158    /// Resource indicator (MCP server URL).
159    resource: String,
160    /// HTTP client for OIDC requests.
161    http_client: HttpClient,
162    /// Secure vault for storing tokens.
163    vault: Vault,
164    /// Internal channel to communicate the auth URL.
165    internal_url_tx: Arc<tokio::sync::Mutex<Option<oneshot::Sender<String>>>>,
166    /// Internal channel to communicate the callback server address.
167    internal_callback_tx: Arc<tokio::sync::Mutex<Option<oneshot::Sender<SocketAddr>>>>,
168    /// Success template HTML.
169    success_html: Arc<String>,
170    /// Failure template HTML.
171    failure_html: Arc<String>,
172    /// Human-friendly name of the identity provider.
173    issuer_name: String,
174    /// Human-friendly name of the protected resource.
175    resource_name: String,
176    /// Timeouts for the interactive flow.
177    timeouts: Timeouts,
178    /// The authorization server's issuer identifier, when discovered.
179    issuer: Option<String>,
180    /// Whether the AS always returns `iss` in the authorization response (RFC 9207).
181    iss_required: bool,
182    /// Latest DPoP nonce provided by the authorization server (RFC 9449 §8).
183    as_nonce: Arc<std::sync::Mutex<Option<String>>>,
184    /// `scopes_supported` of the authorization server.
185    as_scopes: Option<Vec<String>>,
186    /// `scopes_supported` of the protected resource (RFC 9728).
187    resource_scopes: Option<Vec<String>>,
188    /// See [`OidcConfig::request_offline_access`].
189    request_offline_access: bool,
190}
191
192/// Query parameters of the authorization response (RFC 6749 §4.1.2, RFC 9207).
193#[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/// Successful token endpoint response (RFC 6749 §5.1).
208#[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    /// Resolves the authorization server endpoints for `resource`.
218    ///
219    /// `resource_metadata_url` is the (already validated) RFC 9728 metadata URL
220    /// taken from a `WWW-Authenticate` challenge, if any.
221    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        // Precedence: explicit endpoint overrides, then an explicitly configured
238        // discovery URL, then dynamic discovery from the resource (RFC 9728).
239        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            // MCP clients must confirm PKCE support before starting a flow.
274            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    /// Delivers the next authorization URL to `tx` instead of opening a
383    /// browser (for automation and tests).
384    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    /// Reports the address the next loopback callback server binds to.
390    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    /// Full re-authentication flow: PAR -> Loopback Callback -> Token Exchange
396    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        // 1. Generate a new ephemeral DPoP key. It is only stored once the token
408        // exchange succeeds, so a failed flow leaves the current credentials intact.
409        let dpop_key = DpopKey::generate();
410
411        // 2. Prepare PKCE and State
412        let pkce_verifier = random_urlsafe(32);
413        let pkce_challenge = pkce_s256_challenge(&pkce_verifier);
414        let state_val = random_urlsafe(16);
415
416        // 3. Setup Loopback Server to catch the callback
417        let expected_state = state_val.clone();
418        let (server, rx) = self.setup_loopback_server(expected_state).await?;
419
420        // 4. Pushed Authorization Request (PAR)
421        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        // 5. Direct user to Auth URL
439        if let Err(e) = self.open_auth_url(&par_data, url_tx).await {
440            server.stop().await;
441            return Err(e);
442        }
443
444        // 6. Wait for code from callback
445        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        // 7. Token Exchange with DPoP
454        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            // The browser must not keep a connection to this short-lived server:
472            // a later login would otherwise reach the previous, stale handler.
473            .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        // Notify of bound address if requested (for tests)
525        {
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            // Binds the authorization code to our DPoP key (RFC 9449 §10).
568            ("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        // Attempt to open the browser automatically (skip if in tests or explicitly requested)
617        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        // Send the URL to a possible listener (for tests)
644        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    /// Exchanges the authorization code for a DPoP-bound token and stores it,
656    /// together with its key and refresh token.
657    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(&params, 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            // A refresh token we still hold is bound to the previous key.
682            self.vault.delete_refresh_token(user_id)?;
683        }
684        info!("Successfully acquired and stored DPoP-bound token.");
685        Ok(())
686    }
687
688    /// Renews the access token with the stored refresh token (RFC 6749 §6),
689    /// without user interaction.
690    ///
691    /// Returns `Ok(false)` when there is nothing to refresh with or the
692    /// authorization server rejects the refresh token (which is then deleted),
693    /// so the caller can fall back to the interactive flow.
694    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        // Never send a refresh token to an authorization server that didn't issue it.
699        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        // Refresh tokens of public clients are bound to the DPoP key (RFC 9449 §5).
704        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(&params, &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    /// POSTs to the token endpoint with a DPoP proof. The outer error is a
732    /// transport failure; the inner one is the endpoint's error response.
733    ///
734    /// If the AS demands a DPoP nonce (`use_dpop_nonce`), the request is
735    /// retried once with the nonce it supplied (RFC 9449 §8).
736    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    /// Stores the tokens and what they were issued for. `scopes` is `None` for
797    /// a refresh, which keeps the previously requested scopes.
798    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        // A refresh response without refresh_token keeps the current one (RFC 6749 §6).
806        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    /// The stored credential metadata, if the credentials come from this
829    /// authorization server (or their issuer is unknown).
830    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    /// Discards stored credentials issued by a different authorization server
841    /// than the one discovered now. Credentials are bound to their issuer.
842    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    /// Scopes for a new authorization request (MCP scope selection strategy):
859    /// the challenged scopes, else the resource's `scopes_supported`, plus the
860    /// previously requested ones (step-up keeps earlier permissions). `openid`
861    /// is added when the authorization server offers it, `offline_access` too
862    /// when enabled.
863    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        // `openid` keeps working with providers that don't list their scopes.
875        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    /// Retrieves the current access token for a user from the vault.
890    pub fn get_token(&self, user_id: &str) -> Result<Option<String>> {
891        self.vault.get_token(user_id)
892    }
893}
894
895/// The loopback server receiving the authorization response.
896struct LoopbackServer {
897    handle: tokio::task::JoinHandle<()>,
898    shutdown: Option<oneshot::Sender<()>>,
899}
900
901impl LoopbackServer {
902    /// Stops accepting, closes idle connections and frees the port.
903    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
922/// What the loopback callback delivers: the authorization code, or why it failed.
923type 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    /// Issuer the `iss` response parameter must match (RFC 9207).
934    expected_issuer: Option<String>,
935    /// Reject responses without `iss` (the AS advertised support for it).
936    iss_required: bool,
937}
938
939/// Checks an authorization response whose `state` already matched.
940fn evaluate_callback(query: &AuthCallback, state: &AuthServerState) -> CallbackResult {
941    // RFC 9207: a mismatching `iss` means the response comes from another AS
942    // (mix-up attack), so it is checked before anything else is trusted.
943    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    // Requests without the right state are not ours: ignore them without
981    // ending the flow, so a stray page can't cancel the login.
982    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
1032/// Builds the authorization URL for a PAR `request_uri` (RFC 9126 §4),
1033/// keeping any query the endpoint already has.
1034fn 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
1051/// Whether the client id looks like a URL (a Client ID Metadata Document).
1052fn is_url_client_id(client_id: &str) -> bool {
1053    client_id.starts_with("https://") || client_id.starts_with("http://")
1054}
1055
1056/// A Client ID Metadata Document URL must be https and have a path.
1057fn 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
1070/// Returns `len` bytes from the OS CSPRNG, base64url-encoded without padding.
1071fn 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
1077/// PKCE S256 code challenge (RFC 7636 §4.2).
1078fn 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        // RFC 7636 Appendix B test vector
1105        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); // RFC 7636 verifier length for 32 bytes
1115        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>"), "&lt;script&gt;");
1124        assert_eq!(escape_html("a & b"), "a &amp; b");
1125        assert_eq!(
1126            escape_html("\"double quotes\""),
1127            "&quot;double quotes&quot;"
1128        );
1129        assert_eq!(escape_html("'single quotes'"), "&#x27;single quotes&#x27;");
1130        assert_eq!(
1131            escape_html("<img src=x onerror=alert(1)>"),
1132            "&lt;img src=x onerror=alert(1)&gt;"
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        // 1. Success case
1222        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        // 2. Invalid state case
1238        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        // 3. Already authenticated case (tx taken)
1253        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        // 1. Invalid state case (triggering failure template)
1310        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("&lt;script&gt;alert(&#x27;xss&#x27;)&lt;&#x2F;script&gt;"));
1324        assert!(html_err.contains("&lt;b&gt;Bold Resource&lt;&#x2F;b&gt;"));
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 &lt;denied&gt;"));
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        // The flow is still waiting for the real response.
1388        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        // A browser keeps connections alive; the second login's callback must
1458        // reach the second server, not the previous one.
1459        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(); // pools connections like a browser
1487        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        // A PAR failure must not replace the stored DPoP key.
1538        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        // The insecure endpoint is accepted when explicitly allowed (the resource
1595        // name lookup fails quietly against the unreachable host).
1596        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        // The challenge's resource_metadata (a 404 here) must not override the
1712        // explicitly configured discovery URL.
1713        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    /// Serves `doc` as a configured discovery document and runs `discover`.
1726    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        // Challenge scopes win; openid and offline_access are offered by the AS.
1842        assert_eq!(
1843            am.select_scopes(Some(v(&["files:read"])), &[]),
1844            v(&["files:read", "openid", "offline_access"])
1845        );
1846        // Step-up keeps what was requested before.
1847        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        // Without a challenge: the resource's scopes_supported.
1852        am.resource_scopes = Some(v(&["mcp:read"]));
1853        assert_eq!(
1854            am.select_scopes(None, &[]),
1855            v(&["mcp:read", "openid", "offline_access"])
1856        );
1857        // An AS that lists scopes without openid doesn't get it.
1858        am.as_scopes = Some(v(&["mcp:read"]));
1859        assert_eq!(am.select_scopes(None, &[]), v(&["mcp:read"]));
1860        // offline_access is only requested when enabled.
1861        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        // An AS that lists no scopes keeps the historical openid.
1865        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        // Credentials from another AS are discarded.
1880        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        // ...and their refresh token is never sent anywhere.
1888        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        // Credentials from this AS, or without metadata, are kept.
1894        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        // This should fail after 5 retries because the port is occupied by 'listener'
1992        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); // Release port so AuthManager can bind to it
2011
2012        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        // Mock PAR response
2037        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}