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#[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
69pub trait McpOAuthRefreshGuard: Send {}
71
72impl<T: Send> McpOAuthRefreshGuard for T {}
73
74impl McpOAuthCredentials {
75 #[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 #[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 #[must_use]
97 pub fn issuer(mut self, issuer: impl Into<String>) -> Self {
98 self.issuer = Some(issuer.into());
99 self
100 }
101
102 #[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 #[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 #[must_use]
118 pub fn client_id(&self) -> &str {
119 &self.client_id
120 }
121
122 #[must_use]
124 pub fn access_token(&self) -> &str {
125 &self.access_token
126 }
127
128 #[must_use]
130 pub fn refresh_token_value(&self) -> Option<&str> {
131 self.refresh_token.as_deref()
132 }
133
134 #[must_use]
136 pub fn authorization_issuer(&self) -> Option<&str> {
137 self.issuer.as_deref()
138 }
139
140 #[must_use]
142 pub const fn expires_at(&self) -> Option<u64> {
143 self.expires_at_millis
144 }
145
146 #[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#[async_trait]
211pub trait McpOAuthStore: Send + Sync {
212 async fn load(
214 &self,
215 server_name: &str,
216 server_url: &str,
217 ) -> Result<Option<McpOAuthCredentials>, String>;
218
219 async fn save(
221 &self,
222 server_name: &str,
223 server_url: &str,
224 credentials: &McpOAuthCredentials,
225 ) -> Result<(), String>;
226
227 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 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 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(¤t, Some("https://issuer.example")).unwrap();
1043 assert!(validate_refresh_token_issuer(¤t, None).is_err());
1044 }
1045}