Skip to main content

nanocodex_tools/mcp/
oauth.rs

1use std::{
2    collections::{BTreeMap, HashMap},
3    sync::{Arc, OnceLock, Weak},
4    time::{Duration, SystemTime, UNIX_EPOCH},
5};
6
7use async_trait::async_trait;
8use http::{HeaderName, HeaderValue};
9use oauth2::{AccessToken, RefreshToken, Scope, TokenResponse, basic::BasicTokenType};
10use rmcp::transport::{
11    AuthorizationManager, AuthorizationRequest, AuthorizationSession,
12    auth::{
13        AuthClient, AuthorizationMetadata, CredentialStore, InMemoryCredentialStore,
14        OAuthTokenResponse, StoredCredentials, VendorExtraTokenFields,
15    },
16};
17use serde_json::Value;
18use tokio::{
19    io::{AsyncReadExt, AsyncWriteExt},
20    net::TcpListener,
21    sync::{Mutex, RwLock},
22    task::JoinHandle,
23};
24use tracing::{Instrument, info_span};
25
26use super::config::SecretSource;
27
28mod refresh;
29
30const LOGIN_TIMEOUT: Duration = Duration::from_mins(5);
31const MAX_CALLBACK_BYTES: usize = 16 * 1024;
32
33#[derive(Default)]
34pub(crate) struct OAuthMetadataCache {
35    entries: RwLock<HashMap<(String, String), AuthorizationMetadata>>,
36}
37
38impl OAuthMetadataCache {
39    async fn get(&self, server_name: &str, server_url: &str) -> Option<AuthorizationMetadata> {
40        self.entries
41            .read()
42            .await
43            .get(&(server_name.to_owned(), server_url.to_owned()))
44            .cloned()
45    }
46
47    async fn insert(&self, server_name: &str, server_url: &str, metadata: AuthorizationMetadata) {
48        self.entries
49            .write()
50            .await
51            .insert((server_name.to_owned(), server_url.to_owned()), metadata);
52    }
53}
54
55/// OAuth credentials for one Streamable HTTP MCP server.
56///
57/// This value intentionally does not implement `Debug`: access and refresh tokens must not be
58/// emitted by diagnostics. Embedders normally provide these through an [`McpOAuthStore`].
59#[derive(Clone, PartialEq, Eq)]
60pub struct McpOAuthCredentials {
61    client_id: String,
62    access_token: String,
63    refresh_token: Option<String>,
64    issuer: Option<String>,
65    expires_at_millis: Option<u64>,
66    scopes: Vec<String>,
67}
68
69/// An acquired refresh-transaction lock held until its boxed value is dropped.
70pub trait McpOAuthRefreshGuard: Send {}
71
72impl<T: Send> McpOAuthRefreshGuard for T {}
73
74impl McpOAuthCredentials {
75    /// Creates credentials from a dynamically registered client and access token.
76    #[must_use]
77    pub fn new(client_id: impl Into<String>, access_token: impl Into<String>) -> Self {
78        Self {
79            client_id: client_id.into(),
80            access_token: access_token.into(),
81            refresh_token: None,
82            issuer: None,
83            expires_at_millis: None,
84            scopes: Vec::new(),
85        }
86    }
87
88    /// Attaches the optional refresh token.
89    #[must_use]
90    pub fn refresh_token(mut self, refresh_token: impl Into<String>) -> Self {
91        self.refresh_token = Some(refresh_token.into());
92        self
93    }
94
95    /// Binds these credentials to the authorization server that issued them.
96    #[must_use]
97    pub fn issuer(mut self, issuer: impl Into<String>) -> Self {
98        self.issuer = Some(issuer.into());
99        self
100    }
101
102    /// Sets the access-token expiry as Unix epoch milliseconds.
103    #[must_use]
104    pub const fn expires_at_millis(mut self, expires_at_millis: u64) -> Self {
105        self.expires_at_millis = Some(expires_at_millis);
106        self
107    }
108
109    /// Records the scopes granted by the authorization server.
110    #[must_use]
111    pub fn scopes(mut self, scopes: impl IntoIterator<Item = impl Into<String>>) -> Self {
112        self.scopes = scopes.into_iter().map(Into::into).collect();
113        self
114    }
115
116    /// Returns the dynamically registered OAuth client ID.
117    #[must_use]
118    pub fn client_id(&self) -> &str {
119        &self.client_id
120    }
121
122    /// Returns the bearer access token.
123    #[must_use]
124    pub fn access_token(&self) -> &str {
125        &self.access_token
126    }
127
128    /// Returns the refresh token when one was issued.
129    #[must_use]
130    pub fn refresh_token_value(&self) -> Option<&str> {
131        self.refresh_token.as_deref()
132    }
133
134    /// Returns the authorization server issuer bound to these credentials.
135    #[must_use]
136    pub fn authorization_issuer(&self) -> Option<&str> {
137        self.issuer.as_deref()
138    }
139
140    /// Returns access-token expiry as Unix epoch milliseconds.
141    #[must_use]
142    pub const fn expires_at(&self) -> Option<u64> {
143        self.expires_at_millis
144    }
145
146    /// Returns the scopes granted by the authorization server.
147    #[must_use]
148    pub fn granted_scopes(&self) -> &[String] {
149        &self.scopes
150    }
151
152    fn to_token_response(&self) -> OAuthTokenResponse {
153        let mut response = OAuthTokenResponse::new(
154            AccessToken::new(self.access_token.clone()),
155            BasicTokenType::Bearer,
156            VendorExtraTokenFields::default(),
157        );
158        if let Some(refresh_token) = &self.refresh_token {
159            response.set_refresh_token(Some(RefreshToken::new(refresh_token.clone())));
160        }
161        if !self.scopes.is_empty() {
162            response.set_scopes(Some(self.scopes.iter().cloned().map(Scope::new).collect()));
163        }
164        if let Some(expires_at) = self.expires_at_millis {
165            response.set_expires_in(Some(&Duration::from_millis(
166                expires_at.saturating_sub(now_millis()),
167            )));
168        }
169        response
170    }
171
172    fn from_token_response(
173        client_id: String,
174        response: &OAuthTokenResponse,
175        issuer: Option<String>,
176    ) -> Self {
177        let expires_at_millis = response.expires_in().and_then(|expires_in| {
178            now_millis().checked_add(u64::try_from(expires_in.as_millis()).ok()?)
179        });
180        Self {
181            client_id,
182            access_token: response.access_token().secret().to_owned(),
183            refresh_token: response
184                .refresh_token()
185                .map(|token| token.secret().to_owned()),
186            issuer,
187            expires_at_millis,
188            scopes: response
189                .scopes()
190                .map(|scopes| {
191                    scopes
192                        .iter()
193                        .map(|scope| scope.as_ref().to_owned())
194                        .collect()
195                })
196                .unwrap_or_default(),
197        }
198    }
199
200    fn same_token(&self, other: &Self) -> bool {
201        self.client_id == other.client_id
202            && self.access_token == other.access_token
203            && self.refresh_token == other.refresh_token
204            && self.issuer == other.issuer
205            && self.scopes == other.scopes
206    }
207}
208
209/// Persistence selected by an embedding application for MCP OAuth credentials.
210#[async_trait]
211pub trait McpOAuthStore: Send + Sync {
212    /// Loads credentials for one configured server and exact URL.
213    async fn load(
214        &self,
215        server_name: &str,
216        server_url: &str,
217    ) -> Result<Option<McpOAuthCredentials>, String>;
218
219    /// Atomically persists the latest credentials after login or refresh.
220    async fn save(
221        &self,
222        server_name: &str,
223        server_url: &str,
224        credentials: &McpOAuthCredentials,
225    ) -> Result<(), String>;
226
227    /// Serializes one credential's authoritative load, provider refresh, and save transaction.
228    ///
229    /// The default coordinates runtimes in this process. Stores shared by multiple processes must
230    /// override this with a matching cross-process or provider-backed lock and bound lock waits.
231    async fn acquire_refresh_lock(
232        &self,
233        server_name: &str,
234        server_url: &str,
235    ) -> Result<Box<dyn McpOAuthRefreshGuard>, String> {
236        let key = format!("{server_name}\0{server_url}");
237        let lock = {
238            static LOCKS: OnceLock<
239                std::sync::Mutex<HashMap<String, Weak<tokio::sync::Mutex<()>>>>,
240            > = OnceLock::new();
241            let locks = LOCKS.get_or_init(Default::default);
242            let mut locks = locks
243                .lock()
244                .map_err(|_| "MCP OAuth refresh lock registry was poisoned".to_owned())?;
245            locks.retain(|_, lock| lock.strong_count() > 0);
246            if let Some(lock) = locks.get(&key).and_then(Weak::upgrade) {
247                lock
248            } else {
249                let lock = Arc::new(tokio::sync::Mutex::new(()));
250                locks.insert(key, Arc::downgrade(&lock));
251                lock
252            }
253        };
254        Ok(Box::new(lock.lock_owned().await))
255    }
256}
257
258pub(crate) struct OAuthRuntime {
259    server_name: String,
260    server_url: String,
261    manager: Arc<Mutex<AuthorizationManager>>,
262    store: Arc<dyn McpOAuthStore>,
263    authorization_issuer: Option<String>,
264    last_credentials: Mutex<Option<McpOAuthCredentials>>,
265}
266
267impl OAuthRuntime {
268    pub(crate) fn new(
269        server_name: String,
270        server_url: String,
271        manager: Arc<Mutex<AuthorizationManager>>,
272        store: Arc<dyn McpOAuthStore>,
273        authorization_issuer: Option<String>,
274        credentials: McpOAuthCredentials,
275    ) -> Self {
276        Self {
277            server_name,
278            server_url,
279            manager,
280            store,
281            authorization_issuer,
282            last_credentials: Mutex::new(Some(credentials)),
283        }
284    }
285
286    pub(crate) async fn persist_if_changed(&self, parent: &tracing::Span) -> Result<(), String> {
287        let (client_id, response) = self
288            .manager
289            .lock()
290            .await
291            .get_credentials()
292            .await
293            .map_err(|error| format!("failed to read refreshed OAuth credentials: {error}"))?;
294        let Some(response) = response else {
295            return Err("OAuth transport no longer has credentials".to_owned());
296        };
297        let mut credentials = McpOAuthCredentials::from_token_response(
298            client_id,
299            &response,
300            self.authorization_issuer.clone(),
301        );
302        let mut previous = self.last_credentials.lock().await;
303        if let Some(previous) = previous.as_ref() {
304            if response.refresh_token().is_none() {
305                if validate_refresh_token_issuer(previous, self.authorization_issuer.as_deref())
306                    .is_ok()
307                {
308                    credentials
309                        .refresh_token
310                        .clone_from(&previous.refresh_token);
311                } else if previous.refresh_token.is_some() {
312                    // Do not relabel an unbound refresh token with the current issuer merely
313                    // because RMCP returned the still-usable access token staged without it.
314                    credentials.issuer = None;
315                }
316            }
317            if response.scopes().is_none() {
318                credentials.scopes.clone_from(&previous.scopes);
319            }
320            if credentials.same_token(previous) {
321                credentials.expires_at_millis = previous.expires_at_millis;
322            }
323        }
324        if previous.as_ref() == Some(&credentials) {
325            return Ok(());
326        }
327        let span = info_span!(
328            target: "nanocodex_tools",
329            parent: parent,
330            "mcp.oauth.credentials_save",
331            otel.kind = "internal",
332            otel.status_code = tracing::field::Empty,
333            reason = "refresh",
334            status = tracing::field::Empty,
335        );
336        let result = self
337            .store
338            .save(&self.server_name, &self.server_url, &credentials)
339            .instrument(span.clone())
340            .await;
341        span.record(
342            "status",
343            if result.is_ok() {
344                "completed"
345            } else {
346                "failed"
347            },
348        );
349        span.record(
350            "otel.status_code",
351            if result.is_ok() { "OK" } else { "ERROR" },
352        );
353        result?;
354        *previous = Some(credentials);
355        Ok(())
356    }
357}
358
359pub(crate) struct OAuthTransport {
360    pub(crate) client: AuthClient<reqwest::Client>,
361    pub(crate) runtime: Arc<OAuthRuntime>,
362    pub(crate) metadata_cache_hit: bool,
363}
364
365fn credentials_for_manager(
366    credentials: &McpOAuthCredentials,
367    authorization_issuer: Option<&str>,
368) -> McpOAuthCredentials {
369    let mut staged = credentials.clone();
370    if validate_refresh_token_issuer(credentials, authorization_issuer).is_err() {
371        // The access token is still useful at the MCP resource. Never expose an unbound refresh
372        // token to RMCP, which could otherwise send it automatically as the access token expires.
373        staged.refresh_token = None;
374        staged.issuer = None;
375    }
376    staged
377}
378
379fn validate_refresh_token_issuer(
380    credentials: &McpOAuthCredentials,
381    authorization_issuer: Option<&str>,
382) -> Result<(), String> {
383    if credentials.refresh_token.is_none() {
384        return Ok(());
385    }
386    let Some(stored_issuer) = credentials.issuer.as_deref() else {
387        return Err("OAuth refresh credentials are missing an authorization server issuer; authorization required".to_owned());
388    };
389    let Some(authorization_issuer) = authorization_issuer else {
390        return Err(
391            "OAuth metadata did not include an authorization server issuer; authorization required"
392                .to_owned(),
393        );
394    };
395    if stored_issuer != authorization_issuer {
396        return Err("OAuth authorization server issuer changed; authorization required".to_owned());
397    }
398    Ok(())
399}
400
401fn authorization_issuer(metadata: &AuthorizationMetadata) -> Result<Option<String>, String> {
402    metadata
403        .issuer
404        .as_deref()
405        .filter(|issuer| !issuer.trim().is_empty())
406        .map(|issuer| {
407            url::Url::parse(issuer).map_err(|error| {
408                format!("OAuth authorization server issuer is invalid: {error}")
409            })?;
410            Ok(issuer.to_owned())
411        })
412        .transpose()
413}
414
415fn validate_authorization_server_endpoints(metadata: &AuthorizationMetadata) -> Result<(), String> {
416    let authorization_endpoint = url::Url::parse(&metadata.authorization_endpoint)
417        .map_err(|error| format!("OAuth authorization endpoint is invalid: {error}"))?;
418    let token_endpoint = url::Url::parse(&metadata.token_endpoint)
419        .map_err(|error| format!("OAuth token endpoint is invalid: {error}"))?;
420    let issuer = metadata
421        .issuer
422        .as_deref()
423        .filter(|issuer| !issuer.trim().is_empty())
424        .map(url::Url::parse)
425        .transpose()
426        .map_err(|error| format!("OAuth authorization server issuer is invalid: {error}"))?;
427    let issuer_bound_callbacks = metadata
428        .additional_fields
429        .get("authorization_response_iss_parameter_supported")
430        .and_then(Value::as_bool)
431        .unwrap_or(false);
432
433    if issuer_bound_callbacks {
434        if issuer.is_none() {
435            return Err(
436                "OAuth issuer-bound callbacks require an authorization server issuer".to_owned(),
437            );
438        }
439        return Ok(());
440    }
441
442    if let Some(issuer) = issuer {
443        let compatible_provider = matches!(
444            (
445                issuer.as_str(),
446                authorization_endpoint
447                    .origin()
448                    .ascii_serialization()
449                    .as_str(),
450                token_endpoint.origin().ascii_serialization().as_str(),
451            ),
452            (
453                "https://api.figma.com/",
454                "https://www.figma.com",
455                "https://api.figma.com",
456            ) | (
457                "https://agent.robinhood.com/mcp/trading",
458                "https://robinhood.com",
459                "https://api.robinhood.com",
460            )
461        );
462        if authorization_endpoint.origin() == issuer.origin()
463            || authorization_endpoint.origin() == token_endpoint.origin()
464            || compatible_provider
465        {
466            return Ok(());
467        }
468        return Err(
469            "OAuth authorization endpoint origin does not match the authorization server origin without issuer-bound callbacks".to_owned(),
470        );
471    }
472
473    if token_endpoint.origin() != authorization_endpoint.origin() {
474        return Err(
475            "OAuth token endpoint origin does not match the authorization server origin without issuer-bound callbacks".to_owned(),
476        );
477    }
478    Ok(())
479}
480
481pub(crate) async fn transport_from_credentials(
482    server_name: &str,
483    server_url: &str,
484    http_client: reqwest::Client,
485    store: Arc<dyn McpOAuthStore>,
486    credentials: McpOAuthCredentials,
487    metadata_cache: &OAuthMetadataCache,
488) -> Result<OAuthTransport, String> {
489    let mut manager = AuthorizationManager::new(server_url)
490        .await
491        .map_err(|error| format!("failed to initialize MCP OAuth state: {error}"))?;
492    manager
493        .with_client(http_client.clone())
494        .map_err(|error| format!("failed to configure MCP OAuth HTTP client: {error}"))?;
495    let (metadata, metadata_cache_hit) =
496        if let Some(metadata) = metadata_cache.get(server_name, server_url).await {
497            (metadata, true)
498        } else {
499            let metadata = manager
500                .resolve_metadata()
501                .await
502                .map_err(|error| format!("failed to discover MCP OAuth metadata: {error}"))?
503                .metadata;
504            metadata_cache
505                .insert(server_name, server_url, metadata.clone())
506                .await;
507            (metadata, false)
508        };
509    validate_authorization_server_endpoints(&metadata)?;
510    let authorization_issuer = authorization_issuer(&metadata)?;
511    manager.set_metadata(metadata);
512
513    let staged_credentials = credentials_for_manager(&credentials, authorization_issuer.as_deref());
514    let credential_store = InMemoryCredentialStore::new();
515    credential_store
516        .save(
517            StoredCredentials::new(
518                staged_credentials.client_id.clone(),
519                Some(staged_credentials.to_token_response()),
520                staged_credentials.scopes.clone(),
521                Some(now_seconds()),
522            )
523            .with_issuer(staged_credentials.issuer.clone()),
524        )
525        .await
526        .map_err(|error| format!("failed to stage MCP OAuth credentials: {error}"))?;
527    manager.set_credential_store(credential_store);
528    let restored = manager
529        .initialize_from_store()
530        .await
531        .map_err(|error| format!("failed to restore MCP OAuth credentials: {error}"))?;
532    if !restored {
533        return Err("restored MCP OAuth state was not authorized".to_owned());
534    }
535    let client = AuthClient::new(http_client, manager);
536    let runtime = Arc::new(OAuthRuntime::new(
537        server_name.to_owned(),
538        server_url.to_owned(),
539        Arc::clone(&client.auth_manager),
540        store,
541        authorization_issuer,
542        credentials,
543    ));
544    Ok(OAuthTransport {
545        client,
546        runtime,
547        metadata_cache_hit,
548    })
549}
550
551pub(crate) struct OAuthLoginFlow {
552    pub(crate) authorization_url: String,
553    pub(crate) completion: JoinHandle<Result<(), String>>,
554}
555
556pub(crate) async fn begin_login(
557    server_name: String,
558    server_url: String,
559    headers: BTreeMap<String, SecretSource>,
560    store: Arc<dyn McpOAuthStore>,
561) -> Result<OAuthLoginFlow, String> {
562    let client = oauth_http_client(headers)?;
563    let listener = TcpListener::bind("127.0.0.1:0")
564        .await
565        .map_err(|error| format!("failed to bind MCP OAuth callback: {error}"))?;
566    let address = listener
567        .local_addr()
568        .map_err(|error| format!("failed to inspect MCP OAuth callback: {error}"))?;
569    let redirect_uri = format!("http://{address}/callback");
570    let authorization_span = info_span!(
571        target: "nanocodex_tools",
572        "mcp.oauth.authorization_start",
573        otel.kind = "client",
574        otel.status_code = tracing::field::Empty,
575        status = tracing::field::Empty,
576    );
577    let authorization = async {
578        let mut manager = AuthorizationManager::new(&server_url)
579            .await
580            .map_err(|error| format!("failed to discover MCP OAuth metadata: {error}"))?;
581        manager
582            .with_client(client)
583            .map_err(|error| format!("failed to configure MCP OAuth HTTP client: {error}"))?;
584        let metadata = manager
585            .resolve_metadata()
586            .await
587            .map_err(|error| format!("failed to discover MCP OAuth metadata: {error}"))?
588            .metadata;
589        validate_authorization_server_endpoints(&metadata)?;
590        let authorization_issuer = authorization_issuer(&metadata)?;
591        manager.set_metadata(metadata);
592        let session = AuthorizationSession::new(
593            manager,
594            AuthorizationRequest::new(&redirect_uri).with_client_name("Nanocodex"),
595        )
596        .await
597        .map_err(|(_, error)| format!("failed to start MCP OAuth authorization: {error}"))?;
598        let authorization_url = session.get_authorization_url().to_owned();
599        Ok::<_, String>((session, authorization_url, authorization_issuer))
600    }
601    .instrument(authorization_span.clone())
602    .await;
603    authorization_span.record(
604        "status",
605        if authorization.is_ok() {
606            "completed"
607        } else {
608            "failed"
609        },
610    );
611    authorization_span.record(
612        "otel.status_code",
613        if authorization.is_ok() { "OK" } else { "ERROR" },
614    );
615    let (session, authorization_url, authorization_issuer) = authorization?;
616
617    let parent = tracing::Span::current();
618    let completion = tokio::spawn(
619        complete_login(
620            listener,
621            redirect_uri,
622            session,
623            authorization_issuer,
624            store,
625            server_name,
626            server_url,
627        )
628        .instrument(parent),
629    );
630    Ok(OAuthLoginFlow {
631        authorization_url,
632        completion,
633    })
634}
635
636async fn complete_login(
637    listener: TcpListener,
638    redirect_uri: String,
639    session: AuthorizationSession,
640    authorization_issuer: Option<String>,
641    store: Arc<dyn McpOAuthStore>,
642    server_name: String,
643    server_url: String,
644) -> Result<(), String> {
645    let callback_span = info_span!(
646        target: "nanocodex_tools",
647        "mcp.oauth.callback_wait",
648        otel.kind = "server",
649        otel.status_code = tracing::field::Empty,
650        status = tracing::field::Empty,
651    );
652    let callback =
653        match tokio::time::timeout(LOGIN_TIMEOUT, receive_callback(listener, &redirect_uri))
654            .instrument(callback_span.clone())
655            .await
656        {
657            Ok(callback) => callback,
658            Err(_) => Err("timed out waiting for MCP OAuth callback".to_owned()),
659        };
660    callback_span.record(
661        "status",
662        if callback.is_ok() {
663            "completed"
664        } else {
665            "failed"
666        },
667    );
668    callback_span.record(
669        "otel.status_code",
670        if callback.is_ok() { "OK" } else { "ERROR" },
671    );
672    let callback = callback?;
673    let exchange_span = info_span!(
674        target: "nanocodex_tools",
675        "mcp.oauth.code_exchange",
676        otel.kind = "client",
677        otel.status_code = tracing::field::Empty,
678        status = tracing::field::Empty,
679    );
680    let result = session
681        .handle_callback_url(&callback)
682        .instrument(exchange_span.clone())
683        .await
684        .map_err(|error| format!("failed to exchange MCP OAuth code: {error}"));
685    exchange_span.record(
686        "status",
687        if result.is_ok() {
688            "completed"
689        } else {
690            "failed"
691        },
692    );
693    exchange_span.record(
694        "otel.status_code",
695        if result.is_ok() { "OK" } else { "ERROR" },
696    );
697    result?;
698    let (client_id, response) = session
699        .get_credentials()
700        .await
701        .map_err(|error| format!("failed to read MCP OAuth credentials: {error}"))?;
702    let response =
703        response.ok_or_else(|| "MCP OAuth provider returned no credentials".to_owned())?;
704    let credentials =
705        McpOAuthCredentials::from_token_response(client_id, &response, authorization_issuer);
706    let save_span = info_span!(
707        target: "nanocodex_tools",
708        "mcp.oauth.credentials_save",
709        otel.kind = "internal",
710        otel.status_code = tracing::field::Empty,
711        reason = "login",
712        status = tracing::field::Empty,
713    );
714    let saved = store
715        .save(&server_name, &server_url, &credentials)
716        .instrument(save_span.clone())
717        .await;
718    save_span.record("status", if saved.is_ok() { "completed" } else { "failed" });
719    save_span.record(
720        "otel.status_code",
721        if saved.is_ok() { "OK" } else { "ERROR" },
722    );
723    saved
724}
725
726fn oauth_http_client(headers: BTreeMap<String, SecretSource>) -> Result<reqwest::Client, String> {
727    let mut resolved = reqwest::header::HeaderMap::with_capacity(headers.len());
728    for (name, source) in headers {
729        let name = name
730            .parse::<HeaderName>()
731            .map_err(|error| format!("invalid HTTP header name `{name}`: {error}"))?;
732        let value = source.resolve()?;
733        let mut value = HeaderValue::from_str(&value)
734            .map_err(|error| format!("invalid value for HTTP header `{name}`: {error}"))?;
735        value.set_sensitive(true);
736        resolved.insert(name, value);
737    }
738    let replays_plaintext_proxy_credentials =
739        resolved.contains_key(reqwest::header::PROXY_AUTHORIZATION);
740    nanocodex_oai_api::transport::install_default_rustls_crypto_provider();
741    reqwest::Client::builder()
742        .default_headers(resolved)
743        .pool_max_idle_per_host(0)
744        .redirect(super::same_origin_redirect_policy(
745            replays_plaintext_proxy_credentials,
746        ))
747        .build()
748        .map_err(|error| format!("failed to build MCP OAuth HTTP client: {error}"))
749}
750
751async fn receive_callback(listener: TcpListener, redirect_uri: &str) -> Result<String, String> {
752    let (mut stream, _) = listener
753        .accept()
754        .await
755        .map_err(|error| format!("failed to accept MCP OAuth callback: {error}"))?;
756    let mut bytes = Vec::with_capacity(2048);
757    loop {
758        let mut chunk = [0_u8; 1024];
759        let read = stream
760            .read(&mut chunk)
761            .await
762            .map_err(|error| format!("failed to read MCP OAuth callback: {error}"))?;
763        if read == 0 {
764            break;
765        }
766        bytes.extend_from_slice(&chunk[..read]);
767        if bytes.windows(4).any(|window| window == b"\r\n\r\n") {
768            break;
769        }
770        if bytes.len() > MAX_CALLBACK_BYTES {
771            return Err("MCP OAuth callback headers were too large".to_owned());
772        }
773    }
774    let request = std::str::from_utf8(&bytes)
775        .map_err(|_| "MCP OAuth callback was not valid HTTP".to_owned())?;
776    let target = request
777        .lines()
778        .next()
779        .and_then(|line| line.split_whitespace().nth(1))
780        .ok_or_else(|| "MCP OAuth callback did not contain a request target".to_owned())?;
781    let base = reqwest::Url::parse(redirect_uri)
782        .map_err(|error| format!("invalid MCP OAuth redirect URI: {error}"))?;
783    let callback = base
784        .join(target)
785        .map_err(|error| format!("invalid MCP OAuth callback target: {error}"))?;
786    if callback.path() != base.path() {
787        let _ = respond(&mut stream, 400, "Invalid OAuth callback path").await;
788        return Err("MCP OAuth callback used an unexpected path".to_owned());
789    }
790    respond(
791        &mut stream,
792        200,
793        "Authentication received. You may close this window.",
794    )
795    .await?;
796    Ok(callback.to_string())
797}
798
799async fn respond(
800    stream: &mut tokio::net::TcpStream,
801    status: u16,
802    body: &str,
803) -> Result<(), String> {
804    let reason = if status == 200 { "OK" } else { "Bad Request" };
805    let response = format!(
806        "HTTP/1.1 {status} {reason}\r\nContent-Type: text/plain; charset=utf-8\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
807        body.len()
808    );
809    stream
810        .write_all(response.as_bytes())
811        .await
812        .map_err(|error| format!("failed to answer MCP OAuth callback: {error}"))
813}
814
815fn now_millis() -> u64 {
816    let millis = SystemTime::now()
817        .duration_since(UNIX_EPOCH)
818        .unwrap_or(Duration::ZERO)
819        .as_millis();
820    u64::try_from(millis).unwrap_or(u64::MAX)
821}
822
823fn now_seconds() -> u64 {
824    SystemTime::now()
825        .duration_since(UNIX_EPOCH)
826        .unwrap_or(Duration::ZERO)
827        .as_secs()
828}
829
830#[cfg(test)]
831mod tests {
832    use super::*;
833    use tokio::sync::oneshot;
834
835    #[derive(Default)]
836    struct RecordingStore {
837        current: Mutex<Option<McpOAuthCredentials>>,
838        saved: Mutex<Vec<McpOAuthCredentials>>,
839    }
840
841    impl RecordingStore {
842        fn with_credentials(credentials: McpOAuthCredentials) -> Self {
843            Self {
844                current: Mutex::new(Some(credentials)),
845                saved: Mutex::new(Vec::new()),
846            }
847        }
848    }
849
850    #[async_trait]
851    impl McpOAuthStore for RecordingStore {
852        async fn load(
853            &self,
854            _server_name: &str,
855            _server_url: &str,
856        ) -> Result<Option<McpOAuthCredentials>, String> {
857            Ok(self.current.lock().await.clone())
858        }
859
860        async fn save(
861            &self,
862            _server_name: &str,
863            _server_url: &str,
864            credentials: &McpOAuthCredentials,
865        ) -> Result<(), String> {
866            self.saved.lock().await.push(credentials.clone());
867            *self.current.lock().await = Some(credentials.clone());
868            Ok(())
869        }
870    }
871
872    #[tokio::test]
873    async fn oauth_headers_do_not_follow_cross_origin_redirects() {
874        let target = TcpListener::bind("127.0.0.1:0").await.unwrap();
875        let target_url = format!("http://{}/metadata", target.local_addr().unwrap());
876        let (target_requested, mut target_requested_rx) = oneshot::channel();
877        let target_task = tokio::spawn(async move {
878            let (mut stream, _) = target.accept().await.unwrap();
879            let mut request = vec![0_u8; 4096];
880            let read = stream.read(&mut request).await.unwrap();
881            target_requested
882                .send(String::from_utf8_lossy(&request[..read]).into_owned())
883                .unwrap();
884            stream
885                .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\n{}")
886                .await
887                .unwrap();
888        });
889
890        let redirect = TcpListener::bind("127.0.0.1:0").await.unwrap();
891        let source_url = format!("http://{}/metadata", redirect.local_addr().unwrap());
892        let redirect_task = tokio::spawn(async move {
893            let (mut stream, _) = redirect.accept().await.unwrap();
894            let mut request = vec![0_u8; 4096];
895            let read = stream.read(&mut request).await.unwrap();
896            assert!(String::from_utf8_lossy(&request[..read]).contains("x-api-key: secret"));
897            let response = format!(
898                "HTTP/1.1 302 Found\r\nLocation: {target_url}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
899            );
900            stream.write_all(response.as_bytes()).await.unwrap();
901        });
902
903        let client = oauth_http_client(BTreeMap::from([(
904            "x-api-key".to_owned(),
905            SecretSource::Value("secret".to_owned()),
906        )]))
907        .unwrap();
908        let error = client.get(source_url).send().await.unwrap_err();
909        assert!(error.is_redirect(), "{error}");
910        assert!(matches!(
911            target_requested_rx.try_recv(),
912            Err(oneshot::error::TryRecvError::Empty)
913        ));
914
915        target_task.abort();
916        redirect_task.await.unwrap();
917    }
918
919    #[tokio::test]
920    async fn cached_metadata_preserves_refresh_and_rotated_token_persistence() {
921        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
922        let issuer = format!("http://{}", listener.local_addr().unwrap());
923        let server_url = format!("{issuer}/mcp");
924        let token_endpoint = format!("{issuer}/token");
925        let responder_issuer = issuer.clone();
926        let responder = tokio::spawn(async move {
927            loop {
928                let (mut stream, _) = listener.accept().await.unwrap();
929                let mut request = vec![0_u8; 4096];
930                let read = stream.read(&mut request).await.unwrap();
931                let request = String::from_utf8_lossy(&request[..read]);
932                let first_line = request.lines().next().unwrap_or_default();
933                let (status, body, complete) = match first_line {
934                    line if line.starts_with("GET /mcp ") => {
935                        ("404 Not Found", String::new(), false)
936                    }
937                    line if line.contains("oauth-protected-resource") => (
938                        "200 OK",
939                        format!(
940                            r#"{{"resource":"{responder_issuer}/mcp","authorization_servers":["{responder_issuer}"]}}"#
941                        ),
942                        false,
943                    ),
944                    line if line.contains("oauth-authorization-server")
945                        || line.contains("openid-configuration") =>
946                    {
947                        (
948                            "200 OK",
949                            format!(
950                                r#"{{"authorization_endpoint":"{responder_issuer}/authorize","token_endpoint":"{responder_issuer}/token","issuer":"{responder_issuer}"}}"#
951                            ),
952                            false,
953                        )
954                    }
955                    line if line.starts_with("POST /token ") => (
956                        "200 OK",
957                        r#"{"access_token":"refreshed-access","token_type":"Bearer","expires_in":3600,"refresh_token":"rotated-refresh","scope":"mcp:tools"}"#.to_owned(),
958                        true,
959                    ),
960                    _ => panic!("unexpected OAuth fixture request: {first_line}"),
961                };
962                let response = format!(
963                    "HTTP/1.1 {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
964                    body.len()
965                );
966                stream.write_all(response.as_bytes()).await.unwrap();
967                if complete {
968                    break;
969                }
970            }
971        });
972
973        let server_name = "cached";
974        let metadata: AuthorizationMetadata = serde_json::from_value(serde_json::json!({
975            "authorization_endpoint": format!("{issuer}/authorize"),
976            "token_endpoint": token_endpoint,
977            "issuer": issuer,
978        }))
979        .unwrap();
980        let metadata_cache = OAuthMetadataCache::default();
981        metadata_cache
982            .insert(server_name, &server_url, metadata)
983            .await;
984        let credentials = McpOAuthCredentials::new("client", "expired-access")
985            .refresh_token("refresh-token")
986            .issuer(issuer.clone())
987            .expires_at_millis(0)
988            .scopes(["mcp:tools"]);
989        let store = Arc::new(RecordingStore::with_credentials(credentials.clone()));
990
991        nanocodex_oai_api::transport::install_default_rustls_crypto_provider();
992        let transport = transport_from_credentials(
993            server_name,
994            &server_url,
995            reqwest::Client::new(),
996            store.clone(),
997            credentials,
998            &metadata_cache,
999        )
1000        .await
1001        .unwrap();
1002        assert!(transport.metadata_cache_hit);
1003        transport.runtime.refresh_if_needed().await.unwrap();
1004        responder.await.unwrap();
1005
1006        let saved = store.saved.lock().await;
1007        assert_eq!(saved.len(), 1);
1008        assert_eq!(saved[0].access_token(), "refreshed-access");
1009        assert_eq!(saved[0].refresh_token_value(), Some("rotated-refresh"));
1010        assert_eq!(saved[0].authorization_issuer(), Some(issuer.as_str()));
1011        assert_eq!(saved[0].granted_scopes(), ["mcp:tools"]);
1012    }
1013
1014    #[test]
1015    fn oauth_endpoint_identity_rejects_unbound_delegation() {
1016        let metadata: AuthorizationMetadata = serde_json::from_value(serde_json::json!({
1017            "issuer": "https://issuer.example/tenant",
1018            "authorization_endpoint": "https://login.attacker.example/authorize",
1019            "token_endpoint": "https://issuer.example/token"
1020        }))
1021        .unwrap();
1022        let error = validate_authorization_server_endpoints(&metadata).unwrap_err();
1023        assert!(error.contains("authorization endpoint origin"), "{error}");
1024
1025        let mut issuer_bound = metadata;
1026        issuer_bound.additional_fields.insert(
1027            "authorization_response_iss_parameter_supported".to_owned(),
1028            Value::Bool(true),
1029        );
1030        validate_authorization_server_endpoints(&issuer_bound).unwrap();
1031    }
1032
1033    #[test]
1034    fn refresh_tokens_require_the_pinned_authorization_issuer() {
1035        let missing = McpOAuthCredentials::new("client", "access").refresh_token("refresh");
1036        assert!(validate_refresh_token_issuer(&missing, Some("https://issuer.example")).is_err());
1037
1038        let changed = missing.issuer("https://old.example");
1039        assert!(validate_refresh_token_issuer(&changed, Some("https://issuer.example")).is_err());
1040
1041        let current = changed.issuer("https://issuer.example");
1042        validate_refresh_token_issuer(&current, Some("https://issuer.example")).unwrap();
1043        assert!(validate_refresh_token_issuer(&current, None).is_err());
1044    }
1045}