1pub(crate) mod metadata;
17mod pkce;
18pub mod store;
19
20use std::collections::HashMap;
21use std::time::Duration;
22
23use reqwest::Url;
24use serde::Deserialize;
25use tokio::net::TcpListener;
26
27use metadata::{AuthServerMetadata, ProtectedResourceMetadata};
28use pkce::Pkce;
29pub use store::{AuthStore, ServerAuth};
30
31pub type BrowserOpener = std::sync::Arc<dyn Fn(&str) -> bool + Send + Sync>;
40
41const CALLBACK_TIMEOUT: Duration = Duration::from_secs(300);
43
44const DEFAULT_SCOPES: &str = "openid profile email";
46
47#[derive(Debug, Deserialize)]
49struct TokenResponse {
50 access_token: String,
51 #[serde(default)]
52 refresh_token: Option<String>,
53 #[serde(default)]
54 expires_in: Option<u64>,
55 #[serde(default)]
56 scope: Option<String>,
57}
58
59#[derive(Debug, Deserialize)]
61struct RegistrationResponse {
62 client_id: String,
63}
64
65pub struct OAuthClient {
67 http: reqwest::Client,
68}
69
70impl Default for OAuthClient {
71 fn default() -> Self {
72 Self::new()
73 }
74}
75
76impl OAuthClient {
77 pub fn new() -> Self {
79 let http = reqwest::Client::builder()
80 .connect_timeout(Duration::from_secs(30))
81 .timeout(Duration::from_secs(60))
82 .build()
83 .expect("failed to build reqwest client");
84 Self { http }
85 }
86
87 pub async fn login(
94 &self,
95 mcp_url: &str,
96 headers: &HashMap<String, String>,
97 opener: BrowserOpener,
98 now: u64,
99 reuse_client_id: Option<&str>,
100 ) -> anyhow::Result<ServerAuth> {
101 let mcp = Url::parse(mcp_url)
102 .map_err(|e| anyhow::anyhow!("Invalid MCP server url '{}': {}", mcp_url, e))?;
103
104 let (resource, server_meta) = self.discover(&mcp, headers).await?;
105
106 let listener = TcpListener::bind("127.0.0.1:0")
111 .await
112 .expect("binding an ephemeral loopback port cannot fail");
113 let port = listener
114 .local_addr()
115 .expect("a bound listener always has a local address")
116 .port();
117 let redirect_uri = format!("http://127.0.0.1:{port}/callback");
118
119 let client_id = match reuse_client_id {
120 Some(id) => id.to_string(),
121 None => self.register(&server_meta, &redirect_uri).await?,
122 };
123
124 let pkce = Pkce::generate();
125 let scope = if server_meta.scopes_supported.is_empty() {
126 DEFAULT_SCOPES.to_string()
127 } else {
128 server_meta.scopes_supported.join(" ")
129 };
130 let authorize_endpoint = Url::parse(&server_meta.authorization_endpoint)
133 .expect("the authorization endpoint was parsed during metadata validation");
134 let authorize_url = build_authorize_url(
135 &authorize_endpoint,
136 &client_id,
137 &redirect_uri,
138 &pkce,
139 &scope,
140 &resource,
141 );
142
143 println!("Opening your browser to authorize:\n {authorize_url}");
146 if !(*opener)(authorize_url.as_str()) {
147 println!("(couldn't open a browser automatically - open the link above)");
148 }
149
150 let code = wait_for_callback(listener, &pkce.state, CALLBACK_TIMEOUT).await?;
151
152 let token = self
153 .exchange_code(
154 &server_meta.token_endpoint,
155 &client_id,
156 &redirect_uri,
157 &code,
158 &pkce.verifier,
159 &resource,
160 )
161 .await?;
162
163 Ok(build_server_auth(
164 resource,
165 &server_meta,
166 client_id,
167 token,
168 now,
169 ))
170 }
171
172 pub async fn refresh(&self, auth: &ServerAuth, now: u64) -> anyhow::Result<ServerAuth> {
174 let refresh_token = auth
175 .refresh_token
176 .as_deref()
177 .ok_or_else(|| anyhow::anyhow!("no refresh token available"))?;
178
179 let params = [
180 ("grant_type", "refresh_token"),
181 ("refresh_token", refresh_token),
182 ("client_id", auth.client_id.as_str()),
183 ("resource", auth.resource.as_str()),
184 ];
185 let value = self
186 .post_form(&auth.token_endpoint, ¶ms)
187 .await
188 .map_err(|e| anyhow::anyhow!("token refresh failed: {}", e))?;
189 let token: TokenResponse = serde_json::from_value(value)
190 .map_err(|e| anyhow::anyhow!("could not parse token response: {}", e))?;
191
192 let mut refreshed = auth.clone();
193 refreshed.access_token = token.access_token;
194 if let Some(new_refresh) = token.refresh_token {
197 refreshed.refresh_token = Some(new_refresh);
198 }
199 refreshed.expires_at = expires_at(token.expires_in, now);
200 if let Some(scope) = token.scope {
201 refreshed.scope = scope;
202 }
203 Ok(refreshed)
204 }
205
206 pub async fn authorization_header(
215 &self,
216 server_name: &str,
217 store_path: &std::path::Path,
218 now: u64,
219 ) -> anyhow::Result<Option<(String, String)>> {
220 self.authorization_header_with(server_name, store_path, now, None)
221 .await
222 }
223
224 pub async fn authorization_header_with(
232 &self,
233 server_name: &str,
234 store_path: &std::path::Path,
235 now: u64,
236 credentials: Option<&dyn leviath_core::CredentialStore>,
237 ) -> anyhow::Result<Option<(String, String)>> {
238 let mut store = AuthStore::load_with(store_path, credentials)?;
239 let Some(auth) = store.get(server_name) else {
240 return Ok(None);
241 };
242
243 let token = if auth.is_expired_at(now) {
244 let refreshed = self.refresh(auth, now).await.map_err(|e| {
245 anyhow::anyhow!(
246 "MCP server '{server_name}' token expired and could not be \
247 refreshed ({e}); re-authenticate with `lev mcp login {server_name}`"
248 )
249 })?;
250 let access = refreshed.access_token.clone();
251 store.set(server_name, refreshed);
252 store.save_with(store_path, credentials)?;
253 access
254 } else {
255 auth.access_token.clone()
256 };
257
258 Ok(Some((
259 "Authorization".to_string(),
260 format!("Bearer {token}"),
261 )))
262 }
263
264 async fn discover(
266 &self,
267 mcp: &Url,
268 headers: &HashMap<String, String>,
269 ) -> anyhow::Result<(String, AuthServerMetadata)> {
270 let www_authenticate = self.probe_challenge(mcp, headers).await;
273 let hinted = metadata::resource_metadata_url(www_authenticate.as_deref());
274 let resource_meta_url = match hinted {
281 Some(hint) => {
282 let parsed = Url::parse(&hint)
283 .map_err(|e| anyhow::anyhow!("invalid resource_metadata URL '{hint}': {e}"))?;
284 if !metadata::same_origin(&parsed, mcp) {
285 anyhow::bail!(
286 "MCP server at {mcp} pointed resource_metadata at a different origin \
287 ({parsed}) - refusing to follow it"
288 );
289 }
290 parsed
291 }
292 None => metadata::well_known_resource_url(mcp),
293 };
294 self.require_safe_discovery_url(&resource_meta_url)?;
295
296 let value = self
297 .get_json(resource_meta_url.as_str())
298 .await
299 .map_err(|e| anyhow::anyhow!("failed to fetch resource metadata: {}", e))?;
300 let resource_meta: ProtectedResourceMetadata = serde_json::from_value(value)
301 .map_err(|e| anyhow::anyhow!("failed to parse resource metadata: {}", e))?;
302
303 let issuer = resource_meta
304 .authorization_servers
305 .first()
306 .ok_or_else(|| anyhow::anyhow!("resource metadata names no authorization server"))?;
307 let resource = if resource_meta.resource.is_empty() {
310 mcp.to_string()
311 } else {
312 resource_meta.resource.clone()
313 };
314
315 let server_meta = self.fetch_auth_server_metadata(issuer).await?;
316 Ok((resource, server_meta))
317 }
318
319 fn require_safe_discovery_url(&self, url: &Url) -> anyhow::Result<()> {
321 match metadata::is_safe_discovery_url(url) {
322 true => Ok(()),
323 false => anyhow::bail!(
324 "refusing OAuth discovery over an insecure URL ({url}): the flow carries a \
325 bearer token, so it must use https (http is permitted only on loopback)"
326 ),
327 }
328 }
329
330 async fn fetch_auth_server_metadata(&self, issuer: &str) -> anyhow::Result<AuthServerMetadata> {
338 let mut last_err = None;
339 let issuer_url = Url::parse(issuer)
343 .map_err(|e| anyhow::anyhow!("invalid authorization server issuer '{issuer}': {e}"))?;
344 for url in metadata::auth_server_metadata_urls(&issuer_url) {
345 self.require_safe_discovery_url(&url)?;
346 match self.fetch_one_auth_server_metadata(url.as_str()).await {
347 Ok(meta) => {
348 self.validate_auth_server_metadata(&issuer_url, &meta)?;
349 return Ok(meta);
350 }
351 Err(e) => last_err = Some(e),
352 }
353 }
354 Err(anyhow::anyhow!(
355 "failed to fetch authorization server metadata: {}",
356 last_err.expect("at least one candidate URL is always tried")
357 ))
358 }
359
360 fn validate_auth_server_metadata(
372 &self,
373 issuer_url: &Url,
374 meta: &AuthServerMetadata,
375 ) -> anyhow::Result<()> {
376 let issuer = issuer_url.as_str();
377
378 if !meta.issuer.is_empty() {
381 let claimed = Url::parse(&meta.issuer).map_err(|e| {
382 anyhow::anyhow!("invalid issuer '{}' in metadata: {e}", meta.issuer)
383 })?;
384 if !metadata::same_origin(&claimed, issuer_url) {
385 anyhow::bail!(
386 "authorization server metadata claims issuer '{}' but was fetched for \
387 '{issuer}' - refusing (RFC 8414 §3.3)",
388 meta.issuer
389 );
390 }
391 }
392
393 for (label, endpoint) in [
394 ("authorization_endpoint", &meta.authorization_endpoint),
395 ("token_endpoint", &meta.token_endpoint),
396 ] {
397 let parsed = Url::parse(endpoint)
398 .map_err(|e| anyhow::anyhow!("invalid {label} '{endpoint}': {e}"))?;
399 self.require_safe_discovery_url(&parsed)?;
400 if !metadata::same_origin(&parsed, issuer_url) {
401 anyhow::bail!(
402 "{label} '{endpoint}' is not on the issuer's origin ('{issuer}') - refusing"
403 );
404 }
405 }
406 Ok(())
407 }
408
409 async fn fetch_one_auth_server_metadata(
411 &self,
412 url: &str,
413 ) -> anyhow::Result<AuthServerMetadata> {
414 let value = self.get_json(url).await?;
415 Ok(serde_json::from_value(value)?)
416 }
417
418 async fn probe_challenge(
423 &self,
424 mcp: &Url,
425 headers: &HashMap<String, String>,
426 ) -> Option<String> {
427 let mut request = self.http.post(mcp.clone()).body("{}");
428 for (name, value) in headers {
429 request = request.header(name, value);
430 }
431 let response = request.send().await.ok()?;
432 response
433 .headers()
434 .get(reqwest::header::WWW_AUTHENTICATE)
435 .and_then(|v| v.to_str().ok())
436 .map(str::to_string)
437 }
438
439 async fn register(
441 &self,
442 server_meta: &AuthServerMetadata,
443 redirect_uri: &str,
444 ) -> anyhow::Result<String> {
445 let endpoint = server_meta
446 .registration_endpoint
447 .as_deref()
448 .ok_or_else(|| {
449 anyhow::anyhow!(
450 "authorization server does not support dynamic client registration; \
451 a client id must be configured manually"
452 )
453 })?;
454
455 let body = serde_json::json!({
456 "client_name": "Leviath",
457 "redirect_uris": [redirect_uri],
458 "grant_types": ["authorization_code", "refresh_token"],
459 "response_types": ["code"],
460 "token_endpoint_auth_method": "none",
461 });
462 let response = self
463 .http
464 .post(endpoint)
465 .json(&body)
466 .send()
467 .await
468 .map_err(|e| anyhow::anyhow!("client registration request failed: {}", e))?;
469 if !response.status().is_success() {
470 let status = response.status();
471 let text = response.text().await.unwrap_or_default();
472 return Err(anyhow::anyhow!(
473 "client registration failed with HTTP {}: {}",
474 status,
475 text.trim()
476 ));
477 }
478 let registration: RegistrationResponse = response
479 .json()
480 .await
481 .map_err(|e| anyhow::anyhow!("failed to parse registration response: {}", e))?;
482 Ok(registration.client_id)
483 }
484
485 async fn exchange_code(
487 &self,
488 token_endpoint: &str,
489 client_id: &str,
490 redirect_uri: &str,
491 code: &str,
492 verifier: &str,
493 resource: &str,
494 ) -> anyhow::Result<TokenResponse> {
495 let params = [
496 ("grant_type", "authorization_code"),
497 ("code", code),
498 ("redirect_uri", redirect_uri),
499 ("client_id", client_id),
500 ("code_verifier", verifier),
501 ("resource", resource),
502 ];
503 let value = self
504 .post_form(token_endpoint, ¶ms)
505 .await
506 .map_err(|e| anyhow::anyhow!("token exchange failed: {}", e))?;
507 serde_json::from_value(value)
508 .map_err(|e| anyhow::anyhow!("could not parse token response: {}", e))
509 }
510
511 async fn get_json(&self, url: &str) -> anyhow::Result<serde_json::Value> {
517 let response = self.http.get(url).send().await?;
518 if !response.status().is_success() {
519 anyhow::bail!("HTTP {}", response.status());
520 }
521 Ok(response.json().await?)
522 }
523
524 async fn post_form(
528 &self,
529 url: &str,
530 params: &[(&str, &str)],
531 ) -> anyhow::Result<serde_json::Value> {
532 let response = self.http.post(url).form(params).send().await?;
533 let status = response.status();
534 let body = response.text().await.unwrap_or_default();
535 if !status.is_success() {
536 anyhow::bail!("HTTP {}: {}", status, body.trim());
537 }
538 serde_json::from_str(&body)
539 .map_err(|e| anyhow::anyhow!("could not parse token response: {}", e))
540 }
541}
542
543pub struct StoredTokenRefresher {
549 server_name: String,
550 store_path: std::path::PathBuf,
551 clock: fn() -> u64,
553}
554
555impl StoredTokenRefresher {
556 pub fn new(server_name: impl Into<String>, store_path: std::path::PathBuf) -> Self {
558 Self {
559 server_name: server_name.into(),
560 store_path,
561 clock: system_now_secs,
562 }
563 }
564}
565
566fn system_now_secs() -> u64 {
568 std::time::SystemTime::now()
569 .duration_since(std::time::UNIX_EPOCH)
570 .map(|d| d.as_secs())
571 .unwrap_or(0)
572}
573
574#[async_trait::async_trait]
575impl crate::transport::BearerRefresher for StoredTokenRefresher {
576 async fn refresh(&self) -> anyhow::Result<String> {
577 let mut store = AuthStore::load(&self.store_path)?;
578 let auth = store.get(&self.server_name).ok_or_else(|| {
579 anyhow::anyhow!(
580 "no stored credentials for MCP server '{}'",
581 self.server_name
582 )
583 })?;
584 let refreshed = OAuthClient::new().refresh(auth, (self.clock)()).await?;
585 let value = format!("Bearer {}", refreshed.access_token);
586 store.set(&self.server_name, refreshed);
587 store.save(&self.store_path)?;
588 Ok(value)
589 }
590}
591
592fn build_authorize_url(
598 endpoint: &Url,
599 client_id: &str,
600 redirect_uri: &str,
601 pkce: &Pkce,
602 scope: &str,
603 resource: &str,
604) -> Url {
605 let mut url = endpoint.clone();
606 url.query_pairs_mut()
607 .append_pair("response_type", "code")
608 .append_pair("client_id", client_id)
609 .append_pair("redirect_uri", redirect_uri)
610 .append_pair("code_challenge", &pkce.challenge)
611 .append_pair("code_challenge_method", "S256")
612 .append_pair("state", &pkce.state)
613 .append_pair("scope", scope)
614 .append_pair("resource", resource);
616 url
617}
618
619fn build_server_auth(
621 resource: String,
622 server_meta: &AuthServerMetadata,
623 client_id: String,
624 token: TokenResponse,
625 now: u64,
626) -> ServerAuth {
627 ServerAuth {
628 resource,
629 issuer: server_meta.issuer.clone(),
630 authorization_endpoint: server_meta.authorization_endpoint.clone(),
631 token_endpoint: server_meta.token_endpoint.clone(),
632 client_id,
633 access_token: token.access_token,
634 refresh_token: token.refresh_token,
635 expires_at: expires_at(token.expires_in, now),
636 scope: token.scope.unwrap_or_default(),
637 }
638}
639
640fn expires_at(expires_in: Option<u64>, now: u64) -> u64 {
643 match expires_in {
644 Some(secs) => now.saturating_add(secs),
645 None => 0,
646 }
647}
648
649async fn wait_for_callback(
654 listener: TcpListener,
655 expected_state: &str,
656 timeout: Duration,
657) -> anyhow::Result<String> {
658 let accept = async {
659 loop {
660 let (stream, _) = listener
663 .accept()
664 .await
665 .expect("accepting on a bound loopback listener cannot fail");
666 if let Some(result) = handle_callback_connection(stream, expected_state).await? {
669 return Ok(result);
670 }
671 }
672 };
673
674 match tokio::time::timeout(timeout, accept).await {
675 Ok(result) => result,
676 Err(_) => Err(anyhow::anyhow!(
677 "timed out waiting for browser authorization"
678 )),
679 }
680}
681
682async fn handle_callback_connection(
688 mut stream: tokio::net::TcpStream,
689 expected_state: &str,
690) -> anyhow::Result<Option<String>> {
691 use tokio::io::AsyncReadExt;
692
693 let mut buf = vec![0u8; 8192];
694 let n = stream.read(&mut buf).await.unwrap_or(0);
695 let request = String::from_utf8_lossy(&buf[..n]);
696 let Some(target) = request_target(&request) else {
697 return Ok(None);
698 };
699 if !target.starts_with("/callback") {
700 write_response(&mut stream, "404 Not Found", "Not found.").await;
701 return Ok(None);
702 }
703
704 let params = query_params(target);
705 if let Some(error) = params.get("error") {
706 write_response(&mut stream, "400 Bad Request", "Authorization failed.").await;
707 return Err(anyhow::anyhow!("authorization server returned: {}", error));
708 }
709 match (params.get("code"), params.get("state")) {
710 (Some(code), Some(state)) if leviath_core::constant_time_eq(state, expected_state) => {
715 write_response(
716 &mut stream,
717 "200 OK",
718 "Authorization complete - you can close this tab and return to Leviath.",
719 )
720 .await;
721 Ok(Some(code.clone()))
722 }
723 (_, Some(_)) => {
724 write_response(
726 &mut stream,
727 "400 Bad Request",
728 "Invalid authorization state.",
729 )
730 .await;
731 Err(anyhow::anyhow!("OAuth state mismatch - rejecting callback"))
732 }
733 _ => {
734 write_response(
735 &mut stream,
736 "400 Bad Request",
737 "Missing authorization code.",
738 )
739 .await;
740 Err(anyhow::anyhow!("callback missing code or state"))
741 }
742 }
743}
744
745fn request_target(request: &str) -> Option<&str> {
747 let line = request.lines().next()?;
748 let mut parts = line.split_whitespace();
749 let _method = parts.next()?;
750 parts.next()
751}
752
753fn query_params(target: &str) -> HashMap<String, String> {
755 let query = target.split_once('?').map(|(_, q)| q).unwrap_or("");
756 form_urlencoded::parse(query.as_bytes())
757 .map(|(k, v)| (k.into_owned(), v.into_owned()))
758 .collect()
759}
760
761async fn write_response(stream: &mut tokio::net::TcpStream, status: &str, message: &str) {
763 use tokio::io::AsyncWriteExt;
764 let body = format!("<!doctype html><meta charset=utf-8><p>{message}</p>");
765 let response = format!(
766 "HTTP/1.1 {status}\r\nContent-Type: text/html; charset=utf-8\r\n\
767 Content-Length: {}\r\nConnection: close\r\n\r\n{body}",
768 body.len()
769 );
770 let _ = stream.write_all(response.as_bytes()).await;
771 let _ = stream.flush().await;
772}
773
774#[cfg(test)]
775mod tests {
776 use super::*;
777
778 fn fixed_pkce() -> Pkce {
781 Pkce {
782 verifier: "verifier".to_string(),
783 challenge: "challenge".to_string(),
784 state: "state123".to_string(),
785 }
786 }
787
788 #[test]
789 fn authorize_url_carries_every_required_parameter() {
790 let url = build_authorize_url(
791 &Url::parse("https://auth.example.com/authorize").unwrap(),
792 "client-1",
793 "http://127.0.0.1:5000/callback",
794 &fixed_pkce(),
795 "openid profile",
796 "https://mcp.example.com/mcp",
797 );
798 let params: HashMap<_, _> = url.query_pairs().into_owned().collect();
799 assert_eq!(params["response_type"], "code");
800 assert_eq!(params["client_id"], "client-1");
801 assert_eq!(params["redirect_uri"], "http://127.0.0.1:5000/callback");
802 assert_eq!(params["code_challenge"], "challenge");
803 assert_eq!(params["code_challenge_method"], "S256");
804 assert_eq!(params["state"], "state123");
805 assert_eq!(params["scope"], "openid profile");
806 assert_eq!(params["resource"], "https://mcp.example.com/mcp");
808 }
809
810 #[test]
819 fn expires_at_adds_the_relative_lifetime() {
820 assert_eq!(expires_at(Some(3600), 1_000), 4_600);
821 }
822
823 #[test]
824 fn expires_at_is_zero_when_unknown() {
825 assert_eq!(expires_at(None, 1_000), 0);
826 }
827
828 #[test]
831 fn request_target_reads_the_path() {
832 assert_eq!(
833 request_target("GET /callback?code=abc HTTP/1.1\r\nHost: x\r\n\r\n"),
834 Some("/callback?code=abc")
835 );
836 }
837
838 #[test]
839 fn request_target_of_garbage_is_none() {
840 assert_eq!(request_target(""), None);
841 assert_eq!(request_target("GET"), None);
843 assert_eq!(request_target(" \r\n"), None);
845 }
846
847 #[test]
848 fn query_params_parses_pairs() {
849 let params = query_params("/callback?code=abc&state=xyz");
850 assert_eq!(params["code"], "abc");
851 assert_eq!(params["state"], "xyz");
852 }
853
854 #[test]
855 fn query_params_of_a_bare_path_is_empty() {
856 assert!(query_params("/callback").is_empty());
857 }
858
859 fn server_meta() -> AuthServerMetadata {
862 serde_json::from_value(serde_json::json!({
863 "issuer": "https://auth.example.com",
864 "authorization_endpoint": "https://auth.example.com/authorize",
865 "token_endpoint": "https://auth.example.com/token",
866 }))
867 .unwrap()
868 }
869
870 #[test]
871 fn build_server_auth_populates_every_field() {
872 let token = TokenResponse {
873 access_token: "at".to_string(),
874 refresh_token: Some("rt".to_string()),
875 expires_in: Some(3600),
876 scope: Some("openid".to_string()),
877 };
878 let auth = build_server_auth(
879 "https://mcp.example.com/mcp".to_string(),
880 &server_meta(),
881 "client-1".to_string(),
882 token,
883 1_000,
884 );
885 assert_eq!(auth.resource, "https://mcp.example.com/mcp");
886 assert_eq!(auth.issuer, "https://auth.example.com");
887 assert_eq!(auth.client_id, "client-1");
888 assert_eq!(auth.access_token, "at");
889 assert_eq!(auth.refresh_token.as_deref(), Some("rt"));
890 assert_eq!(auth.expires_at, 4_600);
891 assert_eq!(auth.scope, "openid");
892 }
893
894 #[test]
895 fn build_server_auth_defaults_a_missing_scope() {
896 let token = TokenResponse {
897 access_token: "at".to_string(),
898 refresh_token: None,
899 expires_in: None,
900 scope: None,
901 };
902 let auth = build_server_auth(
903 "https://mcp".to_string(),
904 &server_meta(),
905 "c".to_string(),
906 token,
907 0,
908 );
909 assert_eq!(auth.scope, "");
910 assert_eq!(auth.expires_at, 0);
911 assert!(auth.refresh_token.is_none());
912 }
913
914 use tokio::io::{AsyncReadExt, AsyncWriteExt};
921 use tokio::net::TcpStream;
922
923 async fn hit(addr: std::net::SocketAddr, target: &str) -> String {
925 let mut stream = TcpStream::connect(addr).await.unwrap();
926 let request = format!("GET {target} HTTP/1.1\r\nHost: localhost\r\n\r\n");
927 stream.write_all(request.as_bytes()).await.unwrap();
928 stream.flush().await.unwrap();
929 let mut buf = Vec::new();
930 let _ = stream.read_to_end(&mut buf).await;
931 String::from_utf8_lossy(&buf).into_owned()
932 }
933
934 #[tokio::test]
935 async fn callback_returns_the_code_on_a_matching_state() {
936 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
937 let addr = listener.local_addr().unwrap();
938 let server = tokio::spawn(async move {
939 wait_for_callback(listener, "st8", Duration::from_secs(5)).await
940 });
941
942 let response = hit(addr, "/callback?code=the-code&state=st8").await;
943 assert!(response.contains("200 OK"), "got: {response}");
944 assert!(
945 response.contains("Authorization complete"),
946 "got: {response}"
947 );
948 assert_eq!(server.await.unwrap().unwrap(), "the-code");
949 }
950
951 #[tokio::test]
952 async fn callback_skips_unrelated_requests_then_accepts_the_real_one() {
953 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
954 let addr = listener.local_addr().unwrap();
955 let server = tokio::spawn(async move {
956 wait_for_callback(listener, "st8", Duration::from_secs(5)).await
957 });
958
959 let favicon = hit(addr, "/favicon.ico").await;
961 assert!(favicon.contains("404"), "got: {favicon}");
962 let ok = hit(addr, "/callback?code=c&state=st8").await;
963 assert!(ok.contains("200 OK"));
964 assert_eq!(server.await.unwrap().unwrap(), "c");
965 }
966
967 #[tokio::test]
968 async fn callback_rejects_a_mismatched_state() {
969 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
970 let addr = listener.local_addr().unwrap();
971 let server = tokio::spawn(async move {
972 wait_for_callback(listener, "expected", Duration::from_secs(5)).await
973 });
974
975 let response = hit(addr, "/callback?code=c&state=forged").await;
976 assert!(response.contains("400"), "got: {response}");
977 let err = server.await.unwrap().expect_err("mismatch must fail");
978 assert!(err.to_string().contains("state mismatch"), "got: {err}");
979 }
980
981 #[tokio::test]
982 async fn callback_surfaces_an_oauth_error() {
983 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
984 let addr = listener.local_addr().unwrap();
985 let server =
986 tokio::spawn(
987 async move { wait_for_callback(listener, "s", Duration::from_secs(5)).await },
988 );
989
990 let response = hit(addr, "/callback?error=access_denied").await;
991 assert!(response.contains("400"), "got: {response}");
992 let err = server.await.unwrap().expect_err("error param must fail");
993 assert!(err.to_string().contains("access_denied"), "got: {err}");
994 }
995
996 #[tokio::test]
997 async fn callback_rejects_a_request_missing_code_and_state() {
998 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
999 let addr = listener.local_addr().unwrap();
1000 let server =
1001 tokio::spawn(
1002 async move { wait_for_callback(listener, "s", Duration::from_secs(5)).await },
1003 );
1004
1005 let response = hit(addr, "/callback?nothing=here").await;
1006 assert!(response.contains("400"), "got: {response}");
1007 assert!(server.await.unwrap().is_err());
1008 }
1009
1010 #[tokio::test]
1011 async fn handle_callback_ignores_an_empty_connection() {
1012 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1015 let addr = listener.local_addr().unwrap();
1016 let accept = tokio::spawn(async move {
1017 let (stream, _) = listener.accept().await.unwrap();
1018 handle_callback_connection(stream, "s").await
1019 });
1020 let stream = TcpStream::connect(addr).await.unwrap();
1022 drop(stream);
1023 let outcome = accept
1024 .await
1025 .unwrap()
1026 .expect("empty connection is not an error");
1027 assert!(outcome.is_none(), "an empty connection yields no code");
1028 }
1029
1030 use axum::extract::State;
1033 use axum::http::StatusCode;
1034 use axum::response::IntoResponse;
1035 use axum::routing::{get, post};
1036 use axum::{Json, Router};
1037 use std::sync::Arc;
1038 use std::sync::atomic::{AtomicUsize, Ordering};
1039
1040 #[derive(Clone)]
1041 struct MockAs {
1042 base: String,
1043 registrations: Arc<AtomicUsize>,
1044 }
1045
1046 async fn mock_auth_server(variant: &'static str) -> MockAs {
1050 let registrations = Arc::new(AtomicUsize::new(0));
1051 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1052 let base = format!("http://{}", listener.local_addr().unwrap());
1053 let state = MockAs {
1054 base: base.clone(),
1055 registrations: registrations.clone(),
1056 };
1057
1058 let app_state = (base.clone(), variant, registrations.clone());
1059 let app = Router::new()
1060 .route(
1061 "/mcp",
1062 post(|State((base, variant, _)): State<(String, &'static str, Arc<AtomicUsize>)>| async move {
1063 let hint = if variant == "bad_hint" {
1066 "Bearer resource_metadata=\"not a url\"".to_string()
1070 } else {
1071 format!(
1072 "Bearer resource_metadata=\"{base}/.well-known/oauth-protected-resource\""
1073 )
1074 };
1075 (
1076 StatusCode::UNAUTHORIZED,
1077 [(reqwest::header::WWW_AUTHENTICATE, hint)],
1078 )
1079 }),
1080 )
1081 .route(
1082 "/.well-known/oauth-protected-resource",
1083 get(|State((base, variant, _)): State<(String, &'static str, Arc<AtomicUsize>)>| async move {
1084 if variant == "discover_fails" {
1085 return StatusCode::INTERNAL_SERVER_ERROR.into_response();
1086 }
1087 if variant == "resource_not_object" {
1088 return Json(serde_json::json!("just a string")).into_response();
1090 }
1091 let resource = if variant == "no_resource" {
1092 serde_json::Value::String(String::new())
1093 } else {
1094 serde_json::json!(format!("{base}/mcp"))
1095 };
1096 let servers = match variant {
1097 "no_auth_server" => serde_json::json!([]),
1098 "bad_issuer" => serde_json::json!(["not a url"]),
1099 "http_issuer" => serde_json::json!(["http://auth.example.com"]),
1102 _ => serde_json::json!([base]),
1103 };
1104 Json(serde_json::json!({
1105 "resource": resource,
1106 "authorization_servers": servers,
1107 }))
1108 .into_response()
1109 }),
1110 )
1111 .route(
1112 "/.well-known/oauth-authorization-server",
1113 get(|State((base, variant, _)): State<(String, &'static str, Arc<AtomicUsize>)>| async move {
1114 if variant == "no_rfc8414" || variant == "no_metadata" {
1115 return StatusCode::NOT_FOUND.into_response();
1116 }
1117 if variant == "as_bad_rfc8414" {
1118 return Json(serde_json::json!({ "not": "metadata" })).into_response();
1121 }
1122 as_metadata(&base, variant).into_response()
1123 }),
1124 )
1125 .route(
1126 "/.well-known/openid-configuration",
1127 get(|State((base, variant, _)): State<(String, &'static str, Arc<AtomicUsize>)>| async move {
1128 if variant == "no_metadata" {
1129 return StatusCode::NOT_FOUND.into_response();
1130 }
1131 as_metadata(&base, variant).into_response()
1132 }),
1133 )
1134 .route(
1135 "/register",
1136 post(|State((base, variant, regs)): State<(String, &'static str, Arc<AtomicUsize>)>, _body: String| async move {
1137 regs.fetch_add(1, Ordering::SeqCst);
1138 let _ = base;
1139 if variant == "register_fails" {
1140 return (StatusCode::BAD_REQUEST, "invalid_redirect_uri").into_response();
1141 }
1142 if variant == "register_bad_json" {
1143 return (StatusCode::OK, "not json").into_response();
1144 }
1145 Json(serde_json::json!({ "client_id": "registered-client" })).into_response()
1146 }),
1147 )
1148 .route(
1149 "/token",
1150 post(|State((_, variant, _)): State<(String, &'static str, Arc<AtomicUsize>)>, body: String| async move {
1151 if body.contains("refresh_token=bad") {
1153 return (StatusCode::BAD_REQUEST, "invalid_grant").into_response();
1154 }
1155 if variant == "bad_token_json" {
1156 return (StatusCode::OK, "not json").into_response();
1157 }
1158 if variant == "exchange_fails" && body.contains("authorization_code") {
1159 return (StatusCode::BAD_REQUEST, "invalid_grant").into_response();
1160 }
1161 if variant == "minimal_token" {
1162 return Json(serde_json::json!({ "access_token": "minimal" }))
1163 .into_response();
1164 }
1165 if variant == "token_no_access" {
1166 return Json(serde_json::json!({ "wat": true })).into_response();
1168 }
1169 Json(serde_json::json!({
1170 "access_token": "new-access",
1171 "refresh_token": "new-refresh",
1172 "expires_in": 3600,
1173 "scope": "openid",
1174 }))
1175 .into_response()
1176 }),
1177 )
1178 .with_state(app_state);
1179
1180 tokio::spawn(std::future::IntoFuture::into_future(axum::serve(
1181 listener, app,
1182 )));
1183 state
1184 }
1185
1186 fn as_metadata(base: &str, variant: &'static str) -> Json<serde_json::Value> {
1187 let scopes: Vec<&str> = if variant == "no_scopes" {
1188 vec![]
1189 } else {
1190 vec!["openid", "profile"]
1191 };
1192 let authorize = match variant {
1193 "bad_authorize" => "not a url".to_string(),
1194 "foreign_endpoint" => "https://evil.example.com/authorize".to_string(),
1198 "http_endpoint" => "http://evil.example.com/authorize".to_string(),
1201 _ => format!("{base}/authorize"),
1202 };
1203 let issuer = match variant {
1204 "issuer_mismatch" => "https://someone-else.example.com".to_string(),
1207 "as_unparseable_issuer" => "not a url".to_string(),
1211 "no_issuer_field" => String::new(),
1215 _ => base.to_string(),
1216 };
1217 let mut meta = serde_json::json!({
1218 "issuer": issuer,
1219 "authorization_endpoint": authorize,
1220 "token_endpoint": format!("{base}/token"),
1221 "scopes_supported": scopes,
1222 });
1223 if variant != "no_registration" {
1224 meta["registration_endpoint"] = serde_json::json!(format!("{base}/register"));
1225 }
1226 Json(meta)
1227 }
1228
1229 fn drive_callback(authorize_url: &str, state_override: Option<&str>) {
1235 let url = Url::parse(authorize_url).unwrap();
1236 let params: HashMap<_, _> = url.query_pairs().into_owned().collect();
1237 let redirect = params["redirect_uri"].clone();
1238 let state = state_override
1239 .map(String::from)
1240 .unwrap_or_else(|| params["state"].clone());
1241 tokio::spawn(async move {
1243 let callback = format!("{redirect}?code=auth-code&state={state}");
1244 let _ = reqwest::Client::new().get(&callback).send().await;
1245 });
1246 }
1247
1248 fn auto_consent() -> BrowserOpener {
1250 Arc::new(|authorize_url: &str| {
1251 drive_callback(authorize_url, None);
1252 true
1253 })
1254 }
1255
1256 #[tokio::test]
1257 async fn full_login_round_trip() {
1258 let server = mock_auth_server("default").await;
1259 let auth = OAuthClient::new()
1260 .login(
1261 &format!("{}/mcp", server.base),
1262 &HashMap::new(),
1263 auto_consent(),
1264 1_000,
1265 None,
1266 )
1267 .await
1268 .expect("login should complete");
1269
1270 assert_eq!(auth.access_token, "new-access");
1271 assert_eq!(auth.refresh_token.as_deref(), Some("new-refresh"));
1272 assert_eq!(auth.expires_at, 4_600);
1273 assert_eq!(auth.client_id, "registered-client");
1274 assert_eq!(server.registrations.load(Ordering::SeqCst), 1);
1275 }
1276
1277 #[tokio::test]
1278 async fn login_reuses_a_known_client_id_and_skips_registration() {
1279 let server = mock_auth_server("default").await;
1280 OAuthClient::new()
1281 .login(
1282 &format!("{}/mcp", server.base),
1283 &HashMap::new(),
1284 auto_consent(),
1285 0,
1286 Some("existing-client"),
1287 )
1288 .await
1289 .expect("login should complete");
1290 assert_eq!(
1291 server.registrations.load(Ordering::SeqCst),
1292 0,
1293 "a known client id must not re-register"
1294 );
1295 }
1296
1297 #[tokio::test]
1298 async fn login_falls_back_when_rfc8414_metadata_is_malformed() {
1299 let server = mock_auth_server("as_bad_rfc8414").await;
1302 let auth = OAuthClient::new()
1303 .login(
1304 &format!("{}/mcp", server.base),
1305 &HashMap::new(),
1306 auto_consent(),
1307 0,
1308 None,
1309 )
1310 .await
1311 .expect("openid recovery should work");
1312 assert_eq!(auth.access_token, "new-access");
1313 }
1314
1315 #[tokio::test]
1316 async fn login_falls_back_to_openid_configuration() {
1317 let server = mock_auth_server("no_rfc8414").await;
1319 let auth = OAuthClient::new()
1320 .login(
1321 &format!("{}/mcp", server.base),
1322 &HashMap::new(),
1323 auto_consent(),
1324 0,
1325 None,
1326 )
1327 .await
1328 .expect("openid fallback should work");
1329 assert_eq!(auth.access_token, "new-access");
1330 }
1331
1332 #[tokio::test]
1333 async fn login_fails_when_registration_is_unsupported() {
1334 let server = mock_auth_server("no_registration").await;
1335 let err = OAuthClient::new()
1336 .login(
1337 &format!("{}/mcp", server.base),
1338 &HashMap::new(),
1339 auto_consent(),
1340 0,
1341 None,
1342 )
1343 .await
1344 .expect_err("no registration endpoint and no client id must fail");
1345 assert!(
1346 err.to_string().contains("dynamic client registration"),
1347 "got: {err}"
1348 );
1349 }
1350
1351 #[tokio::test]
1352 async fn refresh_rotates_the_tokens() {
1353 let server = mock_auth_server("default").await;
1354 let auth = ServerAuth {
1355 resource: format!("{}/mcp", server.base),
1356 issuer: server.base.clone(),
1357 authorization_endpoint: format!("{}/authorize", server.base),
1358 token_endpoint: format!("{}/token", server.base),
1359 client_id: "c".to_string(),
1360 access_token: "old".to_string(),
1361 refresh_token: Some("good".to_string()),
1362 expires_at: 500,
1363 scope: String::new(),
1364 };
1365 let refreshed = OAuthClient::new().refresh(&auth, 2_000).await.unwrap();
1366 assert_eq!(refreshed.access_token, "new-access");
1367 assert_eq!(refreshed.refresh_token.as_deref(), Some("new-refresh"));
1368 assert_eq!(refreshed.expires_at, 5_600);
1369 }
1370
1371 #[tokio::test]
1372 async fn refresh_keeps_the_old_token_when_none_is_returned() {
1373 let server = mock_auth_server("minimal_token").await;
1376 let auth = ServerAuth {
1377 token_endpoint: format!("{}/token", server.base),
1378 refresh_token: Some("keep-me".to_string()),
1379 scope: "openid".to_string(),
1380 ..Default::default()
1381 };
1382 let refreshed = OAuthClient::new().refresh(&auth, 0).await.unwrap();
1383 assert_eq!(refreshed.access_token, "minimal");
1384 assert_eq!(refreshed.refresh_token.as_deref(), Some("keep-me"));
1385 assert_eq!(refreshed.scope, "openid");
1386 assert_eq!(refreshed.expires_at, 0);
1387 }
1388
1389 #[tokio::test]
1390 async fn authorization_header_is_none_without_stored_auth() {
1391 let dir = tempfile::tempdir().unwrap();
1392 let store = dir.path().join("mcp-auth.json");
1393 let header = OAuthClient::new()
1394 .authorization_header("unknown", &store, 0)
1395 .await
1396 .unwrap();
1397 assert!(header.is_none());
1398 }
1399
1400 #[tokio::test]
1401 async fn authorization_header_returns_a_fresh_token_unchanged() {
1402 let dir = tempfile::tempdir().unwrap();
1403 let store_path = dir.path().join("mcp-auth.json");
1404 let mut store = AuthStore::default();
1405 store.set(
1406 "srv",
1407 ServerAuth {
1408 access_token: "still-good".to_string(),
1409 expires_at: 10_000,
1410 ..Default::default()
1411 },
1412 );
1413 store.save(&store_path).unwrap();
1414
1415 let header = OAuthClient::new()
1416 .authorization_header("srv", &store_path, 1_000)
1417 .await
1418 .unwrap()
1419 .expect("a stored token yields a header");
1420 assert_eq!(
1421 header,
1422 ("Authorization".to_string(), "Bearer still-good".to_string())
1423 );
1424 }
1425
1426 #[tokio::test]
1427 async fn authorization_header_refreshes_an_expired_token_and_persists_it() {
1428 let server = mock_auth_server("default").await;
1429 let dir = tempfile::tempdir().unwrap();
1430 let store_path = dir.path().join("mcp-auth.json");
1431 let mut store = AuthStore::default();
1432 store.set(
1433 "srv",
1434 ServerAuth {
1435 token_endpoint: format!("{}/token", server.base),
1436 access_token: "expired".to_string(),
1437 refresh_token: Some("good".to_string()),
1438 expires_at: 100,
1439 ..Default::default()
1440 },
1441 );
1442 store.save(&store_path).unwrap();
1443
1444 let header = OAuthClient::new()
1445 .authorization_header("srv", &store_path, 1_000)
1446 .await
1447 .unwrap()
1448 .expect("an expired token is refreshed");
1449 assert_eq!(header.1, "Bearer new-access");
1450 let reloaded = AuthStore::load(&store_path).unwrap();
1452 assert_eq!(reloaded.get("srv").unwrap().access_token, "new-access");
1453 }
1454
1455 #[tokio::test]
1456 async fn authorization_header_names_the_login_command_when_refresh_fails() {
1457 let dir = tempfile::tempdir().unwrap();
1458 let store_path = dir.path().join("mcp-auth.json");
1459 let mut store = AuthStore::default();
1460 store.set(
1461 "srv",
1462 ServerAuth {
1463 token_endpoint: "http://127.0.0.1:1/token".to_string(),
1464 access_token: "expired".to_string(),
1465 refresh_token: Some("good".to_string()),
1466 expires_at: 100,
1467 ..Default::default()
1468 },
1469 );
1470 store.save(&store_path).unwrap();
1471
1472 let err = OAuthClient::new()
1473 .authorization_header("srv", &store_path, 1_000)
1474 .await
1475 .expect_err("a dead refresh must fail");
1476 assert!(err.to_string().contains("lev mcp login srv"), "got: {err}");
1477 }
1478
1479 use crate::transport::BearerRefresher;
1482
1483 fn refresher_at(dir: &std::path::Path) -> StoredTokenRefresher {
1484 StoredTokenRefresher {
1485 server_name: "srv".to_string(),
1486 store_path: dir.join("mcp-auth.json"),
1487 clock: || 2_000,
1488 }
1489 }
1490
1491 #[tokio::test]
1492 async fn stored_refresher_rotates_and_persists_the_token() {
1493 let server = mock_auth_server("default").await;
1494 let dir = tempfile::tempdir().unwrap();
1495 let mut store = AuthStore::default();
1496 store.set(
1497 "srv",
1498 ServerAuth {
1499 token_endpoint: format!("{}/token", server.base),
1500 refresh_token: Some("good".to_string()),
1501 expires_at: 1,
1502 ..Default::default()
1503 },
1504 );
1505 let refresher = refresher_at(dir.path());
1506 store.save(&refresher.store_path).unwrap();
1507
1508 let value = refresher.refresh().await.expect("refresh should succeed");
1509 assert_eq!(value, "Bearer new-access");
1510 let reloaded = AuthStore::load(&refresher.store_path).unwrap();
1512 assert_eq!(reloaded.get("srv").unwrap().access_token, "new-access");
1513 }
1514
1515 #[tokio::test]
1516 async fn stored_refresher_errors_without_stored_credentials() {
1517 let dir = tempfile::tempdir().unwrap();
1518 let refresher = refresher_at(dir.path());
1519 let err = refresher.refresh().await.expect_err("no creds must fail");
1521 assert!(
1522 err.to_string().contains("no stored credentials"),
1523 "got: {err}"
1524 );
1525 }
1526
1527 #[tokio::test]
1528 async fn stored_refresher_surfaces_a_refresh_failure() {
1529 let dir = tempfile::tempdir().unwrap();
1530 let mut store = AuthStore::default();
1531 store.set(
1532 "srv",
1533 ServerAuth {
1534 token_endpoint: "http://127.0.0.1:1/token".to_string(),
1535 refresh_token: Some("good".to_string()),
1536 expires_at: 1,
1537 ..Default::default()
1538 },
1539 );
1540 let refresher = refresher_at(dir.path());
1541 store.save(&refresher.store_path).unwrap();
1542 assert!(refresher.refresh().await.is_err());
1543 }
1544
1545 #[test]
1546 fn stored_refresher_new_uses_the_system_clock() {
1547 let r = StoredTokenRefresher::new("s", std::path::PathBuf::from("/tmp/x"));
1548 assert!((r.clock)() > 1_600_000_000);
1549 }
1550
1551 #[test]
1552 fn system_now_secs_advances_past_the_epoch() {
1553 assert!(system_now_secs() > 1_600_000_000);
1554 }
1555
1556 #[tokio::test]
1557 async fn authorization_header_surfaces_an_unreadable_store() {
1558 let dir = tempfile::tempdir().unwrap();
1560 assert!(
1561 OAuthClient::new()
1562 .authorization_header("srv", dir.path(), 0)
1563 .await
1564 .is_err()
1565 );
1566 }
1567
1568 #[tokio::test]
1569 async fn authorization_header_surfaces_an_unwritable_store_after_refresh() {
1570 let server = mock_auth_server("default").await;
1573 let dir = tempfile::tempdir().unwrap();
1574 let store_path = dir.path().join("mcp-auth.json");
1575 let mut store = AuthStore::default();
1576 store.set(
1577 "srv",
1578 ServerAuth {
1579 token_endpoint: format!("{}/token", server.base),
1580 refresh_token: Some("good".to_string()),
1581 expires_at: 100,
1582 ..Default::default()
1583 },
1584 );
1585 store.save(&store_path).unwrap();
1586 let mut perms = std::fs::metadata(&store_path).unwrap().permissions();
1587 perms.set_readonly(true);
1588 std::fs::set_permissions(&store_path, perms).unwrap();
1589
1590 assert!(
1591 OAuthClient::new()
1592 .authorization_header("srv", &store_path, 1_000)
1593 .await
1594 .is_err()
1595 );
1596 }
1597
1598 #[tokio::test]
1599 async fn refresh_without_a_token_is_an_error() {
1600 let mut auth = ServerAuth {
1601 token_endpoint: "http://127.0.0.1:1/token".to_string(),
1602 ..Default::default()
1603 };
1604 auth.refresh_token = None;
1605 let err = OAuthClient::new()
1606 .refresh(&auth, 0)
1607 .await
1608 .expect_err("no refresh token must fail");
1609 assert!(err.to_string().contains("no refresh token"), "got: {err}");
1610 }
1611
1612 #[tokio::test]
1613 async fn refresh_surfaces_a_rejected_grant() {
1614 let server = mock_auth_server("default").await;
1615 let auth = ServerAuth {
1616 token_endpoint: format!("{}/token", server.base),
1617 refresh_token: Some("bad".to_string()),
1618 ..Default::default()
1619 };
1620 let err = OAuthClient::new()
1621 .refresh(&auth, 0)
1622 .await
1623 .expect_err("a rejected grant must fail");
1624 assert!(err.to_string().contains("refresh failed"), "got: {err}");
1625 }
1626
1627 #[tokio::test]
1628 async fn discovery_without_a_www_authenticate_uses_the_well_known_path() {
1629 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1632 let base = format!("http://{}", listener.local_addr().unwrap());
1633 let app = Router::new()
1634 .route("/mcp", post(|| async { StatusCode::OK }))
1635 .route(
1636 "/.well-known/oauth-protected-resource",
1637 get({
1638 let base = base.clone();
1639 move || {
1640 let base = base.clone();
1641 async move {
1642 Json(serde_json::json!({
1643 "resource": format!("{base}/mcp"),
1644 "authorization_servers": [base],
1645 }))
1646 }
1647 }
1648 }),
1649 )
1650 .route(
1651 "/.well-known/oauth-authorization-server",
1652 get({
1653 let base = base.clone();
1654 move || {
1655 let base = base.clone();
1656 async move { as_metadata(&base, "default") }
1657 }
1658 }),
1659 )
1660 .route(
1661 "/register",
1662 post(|| async { Json(serde_json::json!({ "client_id": "c" })) }),
1663 )
1664 .route(
1665 "/token",
1666 post(|| async {
1667 Json(serde_json::json!({"access_token": "at", "expires_in": 60}))
1668 }),
1669 );
1670 tokio::spawn(std::future::IntoFuture::into_future(axum::serve(
1671 listener, app,
1672 )));
1673
1674 let auth = OAuthClient::new()
1675 .login(
1676 &format!("{base}/mcp"),
1677 &HashMap::new(),
1678 auto_consent(),
1679 0,
1680 None,
1681 )
1682 .await
1683 .expect("well-known discovery should work");
1684 assert_eq!(auth.access_token, "at");
1685 }
1686
1687 #[test]
1688 fn oauth_client_default_matches_new() {
1689 let _ = OAuthClient::default();
1691 }
1692
1693 #[tokio::test]
1694 async fn callback_times_out_when_no_redirect_arrives() {
1695 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1696 let err = wait_for_callback(listener, "s", Duration::from_millis(100))
1698 .await
1699 .expect_err("must time out");
1700 assert!(err.to_string().contains("timed out"), "got: {err}");
1701 }
1702
1703 #[tokio::test]
1704 async fn login_sends_configured_probe_headers() {
1705 let server = mock_auth_server("default").await;
1708 let headers = HashMap::from([("X-Probe".to_string(), "1".to_string())]);
1709 OAuthClient::new()
1710 .login(
1711 &format!("{}/mcp", server.base),
1712 &headers,
1713 auto_consent(),
1714 0,
1715 None,
1716 )
1717 .await
1718 .expect("login with probe headers should complete");
1719 }
1720
1721 #[tokio::test]
1722 async fn login_still_completes_when_the_browser_cannot_open() {
1723 let failing_opener: BrowserOpener = Arc::new(|authorize_url: &str| {
1727 drive_callback(authorize_url, None);
1728 false
1729 });
1730 let server = mock_auth_server("default").await;
1731 OAuthClient::new()
1732 .login(
1733 &format!("{}/mcp", server.base),
1734 &HashMap::new(),
1735 failing_opener,
1736 0,
1737 None,
1738 )
1739 .await
1740 .expect("login should complete even without a browser");
1741 }
1742
1743 #[tokio::test]
1744 async fn login_uses_default_scopes_when_the_server_advertises_none() {
1745 let server = mock_auth_server("no_scopes").await;
1746 OAuthClient::new()
1747 .login(
1748 &format!("{}/mcp", server.base),
1749 &HashMap::new(),
1750 auto_consent(),
1751 0,
1752 None,
1753 )
1754 .await
1755 .expect("login should complete with default scopes");
1756 }
1757
1758 #[tokio::test]
1759 async fn login_falls_back_to_the_mcp_url_when_resource_is_omitted() {
1760 let server = mock_auth_server("no_resource").await;
1761 let auth = OAuthClient::new()
1762 .login(
1763 &format!("{}/mcp", server.base),
1764 &HashMap::new(),
1765 auto_consent(),
1766 0,
1767 None,
1768 )
1769 .await
1770 .expect("login should complete");
1771 assert_eq!(auth.resource, format!("{}/mcp", server.base));
1773 }
1774
1775 #[tokio::test]
1776 async fn login_fails_when_registration_is_rejected() {
1777 let server = mock_auth_server("register_fails").await;
1778 let err = OAuthClient::new()
1779 .login(
1780 &format!("{}/mcp", server.base),
1781 &HashMap::new(),
1782 auto_consent(),
1783 0,
1784 None,
1785 )
1786 .await
1787 .expect_err("a rejected registration must fail");
1788 assert!(
1789 err.to_string().contains("registration failed"),
1790 "got: {err}"
1791 );
1792 }
1793
1794 #[tokio::test]
1795 async fn discovery_fails_when_no_metadata_document_is_reachable() {
1796 let server = mock_auth_server("no_metadata").await;
1798 let err = OAuthClient::new()
1799 .login(
1800 &format!("{}/mcp", server.base),
1801 &HashMap::new(),
1802 auto_consent(),
1803 0,
1804 None,
1805 )
1806 .await
1807 .expect_err("no reachable metadata must fail");
1808 assert!(
1809 err.to_string().contains("authorization server metadata"),
1810 "got: {err}"
1811 );
1812 }
1813
1814 #[tokio::test]
1815 async fn discovery_fails_when_resource_metadata_is_unavailable() {
1816 let server = mock_auth_server("discover_fails").await;
1817 let err = OAuthClient::new()
1818 .login(
1819 &format!("{}/mcp", server.base),
1820 &HashMap::new(),
1821 auto_consent(),
1822 0,
1823 None,
1824 )
1825 .await
1826 .expect_err("a 500 on resource metadata must fail");
1827 assert!(err.to_string().contains("resource metadata"), "got: {err}");
1828 }
1829
1830 #[tokio::test]
1831 async fn login_fails_when_registration_returns_bad_json() {
1832 let server = mock_auth_server("register_bad_json").await;
1833 let err = OAuthClient::new()
1834 .login(
1835 &format!("{}/mcp", server.base),
1836 &HashMap::new(),
1837 auto_consent(),
1838 0,
1839 None,
1840 )
1841 .await
1842 .expect_err("unparseable registration must fail");
1843 assert!(
1844 err.to_string().contains("registration response"),
1845 "got: {err}"
1846 );
1847 }
1848
1849 #[tokio::test]
1850 async fn login_fails_when_the_token_exchange_is_rejected() {
1851 let server = mock_auth_server("exchange_fails").await;
1852 let err = OAuthClient::new()
1853 .login(
1854 &format!("{}/mcp", server.base),
1855 &HashMap::new(),
1856 auto_consent(),
1857 0,
1858 None,
1859 )
1860 .await
1861 .expect_err("a rejected code exchange must fail");
1862 assert!(
1863 err.to_string().contains("token exchange failed"),
1864 "got: {err}"
1865 );
1866 }
1867
1868 #[tokio::test]
1869 async fn login_fails_when_the_token_response_is_not_json() {
1870 let server = mock_auth_server("bad_token_json").await;
1871 let err = OAuthClient::new()
1872 .login(
1873 &format!("{}/mcp", server.base),
1874 &HashMap::new(),
1875 auto_consent(),
1876 0,
1877 None,
1878 )
1879 .await
1880 .expect_err("an unparseable token response must fail");
1881 assert!(
1882 err.to_string().contains("parse token response"),
1883 "got: {err}"
1884 );
1885 }
1886
1887 #[tokio::test]
1893 async fn login_refuses_a_cross_origin_resource_metadata_hint() {
1894 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1897 let base = format!("http://{}", listener.local_addr().unwrap());
1898 let app = axum::Router::new().route(
1899 "/mcp",
1900 post(|| async {
1901 (
1902 StatusCode::UNAUTHORIZED,
1903 [(
1904 reqwest::header::WWW_AUTHENTICATE,
1905 "Bearer resource_metadata=\"http://169.254.169.254/latest/meta-data/\"",
1906 )],
1907 )
1908 }),
1909 );
1910 tokio::spawn(std::future::IntoFuture::into_future(axum::serve(
1911 listener, app,
1912 )));
1913
1914 let err = OAuthClient::new()
1915 .login(
1916 &format!("{base}/mcp"),
1917 &HashMap::new(),
1918 auto_consent(),
1919 0,
1920 None,
1921 )
1922 .await
1923 .expect_err("a cross-origin resource_metadata hint must be refused");
1924 let msg = err.to_string();
1925 assert!(msg.contains("different origin"), "got: {msg}");
1926 assert!(msg.contains("169.254.169.254"), "got: {msg}");
1927 }
1928
1929 #[tokio::test]
1932 async fn login_refuses_a_malformed_resource_metadata_hint() {
1933 let server = mock_auth_server("bad_hint").await;
1934 let err = OAuthClient::new()
1935 .login(
1936 &format!("{}/mcp", server.base),
1937 &HashMap::new(),
1938 auto_consent(),
1939 0,
1940 None,
1941 )
1942 .await
1943 .expect_err("a malformed hint must be refused");
1944 assert!(
1945 err.to_string().contains("invalid resource_metadata URL"),
1946 "got: {err}"
1947 );
1948 }
1949
1950 #[tokio::test]
1955 async fn login_refuses_metadata_claiming_a_different_issuer() {
1956 let server = mock_auth_server("issuer_mismatch").await;
1957 let err = OAuthClient::new()
1958 .login(
1959 &format!("{}/mcp", server.base),
1960 &HashMap::new(),
1961 auto_consent(),
1962 0,
1963 None,
1964 )
1965 .await
1966 .expect_err("an issuer mismatch must be refused");
1967 assert!(err.to_string().contains("RFC 8414"), "got: {err}");
1968 }
1969
1970 #[tokio::test]
1972 async fn login_refuses_metadata_with_an_unparseable_issuer() {
1973 let server = mock_auth_server("as_unparseable_issuer").await;
1974 let err = OAuthClient::new()
1975 .login(
1976 &format!("{}/mcp", server.base),
1977 &HashMap::new(),
1978 auto_consent(),
1979 0,
1980 None,
1981 )
1982 .await
1983 .expect_err("an unparseable issuer must be refused");
1984 assert!(err.to_string().contains("invalid issuer"), "got: {err}");
1985 }
1986
1987 #[tokio::test]
1991 async fn login_refuses_an_endpoint_off_the_issuers_origin() {
1992 let server = mock_auth_server("foreign_endpoint").await;
1993 let err = OAuthClient::new()
1994 .login(
1995 &format!("{}/mcp", server.base),
1996 &HashMap::new(),
1997 auto_consent(),
1998 0,
1999 None,
2000 )
2001 .await
2002 .expect_err("a foreign endpoint must be refused");
2003 assert!(
2004 err.to_string().contains("is not on the issuer's origin"),
2005 "got: {err}"
2006 );
2007 }
2008
2009 #[tokio::test]
2014 async fn metadata_without_an_issuer_field_still_completes() {
2015 let server = mock_auth_server("no_issuer_field").await;
2016 let auth = OAuthClient::new()
2017 .login(
2018 &format!("{}/mcp", server.base),
2019 &HashMap::new(),
2020 auto_consent(),
2021 0,
2022 None,
2023 )
2024 .await
2025 .expect("an absent issuer is not itself a failure");
2026 assert!(!auth.access_token.is_empty());
2027 }
2028
2029 #[tokio::test]
2037 async fn every_step_of_discovery_refuses_remote_http() {
2038 let insecure = |err: anyhow::Error| {
2039 let msg = err.to_string();
2040 assert!(msg.contains("refusing OAuth discovery"), "got: {msg}");
2041 };
2042
2043 insecure(
2046 OAuthClient::new()
2047 .login(
2048 "http://mcp.example.invalid/mcp",
2049 &HashMap::new(),
2050 auto_consent(),
2051 0,
2052 None,
2053 )
2054 .await
2055 .expect_err("a remote http MCP URL must be refused"),
2056 );
2057
2058 let server = mock_auth_server("http_issuer").await;
2060 insecure(
2061 OAuthClient::new()
2062 .login(
2063 &format!("{}/mcp", server.base),
2064 &HashMap::new(),
2065 auto_consent(),
2066 0,
2067 None,
2068 )
2069 .await
2070 .expect_err("a remote http issuer must be refused"),
2071 );
2072
2073 let server = mock_auth_server("http_endpoint").await;
2075 insecure(
2076 OAuthClient::new()
2077 .login(
2078 &format!("{}/mcp", server.base),
2079 &HashMap::new(),
2080 auto_consent(),
2081 0,
2082 None,
2083 )
2084 .await
2085 .expect_err("a remote http endpoint must be refused"),
2086 );
2087 }
2088
2089 #[test]
2094 fn discovery_over_remote_http_is_refused() {
2095 let client = OAuthClient::new();
2096 let err = client
2097 .require_safe_discovery_url(&Url::parse("http://auth.example.com/x").unwrap())
2098 .expect_err("remote http must be refused");
2099 assert!(err.to_string().contains("must use https"), "got: {err}");
2100 assert!(
2101 client
2102 .require_safe_discovery_url(&Url::parse("https://auth.example.com/x").unwrap())
2103 .is_ok()
2104 );
2105 }
2106
2107 #[tokio::test]
2108 async fn login_fails_when_the_authorize_endpoint_is_malformed() {
2109 let server = mock_auth_server("bad_authorize").await;
2110 let err = OAuthClient::new()
2111 .login(
2112 &format!("{}/mcp", server.base),
2113 &HashMap::new(),
2114 auto_consent(),
2115 0,
2116 None,
2117 )
2118 .await
2119 .expect_err("a bad authorize endpoint must fail");
2120 assert!(
2125 err.to_string().contains("authorization_endpoint"),
2126 "got: {err}"
2127 );
2128 }
2129
2130 #[tokio::test]
2131 async fn login_fails_when_no_authorization_server_is_named() {
2132 let server = mock_auth_server("no_auth_server").await;
2133 let err = OAuthClient::new()
2134 .login(
2135 &format!("{}/mcp", server.base),
2136 &HashMap::new(),
2137 auto_consent(),
2138 0,
2139 None,
2140 )
2141 .await
2142 .expect_err("empty authorization_servers must fail");
2143 assert!(
2144 err.to_string().contains("no authorization server"),
2145 "got: {err}"
2146 );
2147 }
2148
2149 #[tokio::test]
2150 async fn login_fails_when_the_issuer_is_malformed() {
2151 let server = mock_auth_server("bad_issuer").await;
2152 let err = OAuthClient::new()
2153 .login(
2154 &format!("{}/mcp", server.base),
2155 &HashMap::new(),
2156 auto_consent(),
2157 0,
2158 None,
2159 )
2160 .await
2161 .expect_err("a bad issuer must fail");
2162 assert!(err.to_string().contains("issuer"), "got: {err}");
2163 }
2164
2165 #[tokio::test]
2166 async fn login_fails_when_the_callback_is_forged() {
2167 let forge: BrowserOpener = Arc::new(|authorize_url: &str| {
2170 drive_callback(authorize_url, Some("WRONG"));
2171 true
2172 });
2173 let server = mock_auth_server("default").await;
2174 let err = OAuthClient::new()
2175 .login(
2176 &format!("{}/mcp", server.base),
2177 &HashMap::new(),
2178 forge,
2179 0,
2180 None,
2181 )
2182 .await
2183 .expect_err("a forged callback must fail login");
2184 assert!(err.to_string().contains("state mismatch"), "got: {err}");
2185 }
2186
2187 #[tokio::test]
2194 async fn get_json_errors_on_a_dead_connection() {
2195 let err = OAuthClient::new()
2196 .get_json("http://127.0.0.1:1/x")
2197 .await
2198 .expect_err("a refused connection must fail");
2199 let _ = err;
2200 }
2201
2202 #[tokio::test]
2203 async fn get_json_errors_on_a_non_success_status() {
2204 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
2205 let base = format!("http://{}", listener.local_addr().unwrap());
2206 let app = Router::new().route("/x", get(|| async { StatusCode::NOT_FOUND }));
2207 tokio::spawn(std::future::IntoFuture::into_future(axum::serve(
2208 listener, app,
2209 )));
2210 let err = OAuthClient::new()
2211 .get_json(&format!("{base}/x"))
2212 .await
2213 .expect_err("404 must fail");
2214 assert!(err.to_string().contains("404"), "got: {err}");
2215 }
2216
2217 #[tokio::test]
2218 async fn get_json_errors_on_an_unparseable_body() {
2219 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
2220 let base = format!("http://{}", listener.local_addr().unwrap());
2221 let app = Router::new().route("/x", get(|| async { "not json" }));
2222 tokio::spawn(std::future::IntoFuture::into_future(axum::serve(
2223 listener, app,
2224 )));
2225 assert!(
2226 OAuthClient::new()
2227 .get_json(&format!("{base}/x"))
2228 .await
2229 .is_err()
2230 );
2231 }
2232
2233 #[tokio::test]
2234 async fn post_form_errors_on_a_dead_connection() {
2235 assert!(
2236 OAuthClient::new()
2237 .post_form("http://127.0.0.1:1/token", &[("a", "b")])
2238 .await
2239 .is_err()
2240 );
2241 }
2242
2243 #[tokio::test]
2244 async fn register_errors_on_a_dead_connection() {
2245 let meta: AuthServerMetadata = serde_json::from_value(serde_json::json!({
2246 "issuer": "https://x",
2247 "authorization_endpoint": "https://x/a",
2248 "token_endpoint": "https://x/t",
2249 "registration_endpoint": "http://127.0.0.1:1/register",
2250 }))
2251 .unwrap();
2252 let err = OAuthClient::new()
2253 .register(&meta, "http://127.0.0.1:5000/callback")
2254 .await
2255 .expect_err("a dead registration endpoint must fail");
2256 assert!(
2257 err.to_string().contains("registration request failed"),
2258 "got: {err}"
2259 );
2260 }
2261
2262 #[tokio::test]
2263 async fn probe_challenge_of_a_dead_server_is_none() {
2264 let mcp = Url::parse("http://127.0.0.1:1/mcp").unwrap();
2265 assert!(
2266 OAuthClient::new()
2267 .probe_challenge(&mcp, &HashMap::new())
2268 .await
2269 .is_none()
2270 );
2271 }
2272
2273 #[tokio::test]
2274 async fn login_fails_when_resource_metadata_is_not_an_object() {
2275 let server = mock_auth_server("resource_not_object").await;
2276 let err = OAuthClient::new()
2277 .login(
2278 &format!("{}/mcp", server.base),
2279 &HashMap::new(),
2280 auto_consent(),
2281 0,
2282 None,
2283 )
2284 .await
2285 .expect_err("malformed resource metadata must fail");
2286 assert!(
2287 err.to_string().contains("parse resource metadata"),
2288 "got: {err}"
2289 );
2290 }
2291
2292 #[tokio::test]
2293 async fn login_fails_when_the_token_lacks_an_access_token() {
2294 let server = mock_auth_server("token_no_access").await;
2295 let err = OAuthClient::new()
2296 .login(
2297 &format!("{}/mcp", server.base),
2298 &HashMap::new(),
2299 auto_consent(),
2300 0,
2301 None,
2302 )
2303 .await
2304 .expect_err("a token without access_token must fail");
2305 assert!(
2306 err.to_string().contains("parse token response"),
2307 "got: {err}"
2308 );
2309 }
2310
2311 #[tokio::test]
2312 async fn refresh_fails_when_the_token_lacks_an_access_token() {
2313 let server = mock_auth_server("token_no_access").await;
2314 let auth = ServerAuth {
2315 token_endpoint: format!("{}/token", server.base),
2316 refresh_token: Some("good".to_string()),
2317 ..Default::default()
2318 };
2319 let err = OAuthClient::new()
2320 .refresh(&auth, 0)
2321 .await
2322 .expect_err("a malformed refresh token response must fail");
2323 assert!(
2324 err.to_string().contains("parse token response"),
2325 "got: {err}"
2326 );
2327 }
2328
2329 #[tokio::test]
2330 async fn login_rejects_a_bad_mcp_url() {
2331 let err = OAuthClient::new()
2332 .login("not a url", &HashMap::new(), auto_consent(), 0, None)
2333 .await
2334 .expect_err("bad url must fail");
2335 assert!(
2336 err.to_string().contains("Invalid MCP server url"),
2337 "got: {err}"
2338 );
2339 }
2340}