返回 CodeWhale
oauth.rs
根目录 / crates / tui / src / mcp / oauth.rs
1 use super::http_client::McpHttpClient;
2 use crate::network_policy::NetworkPolicyDecider;
3 use std::collections::HashMap;
4 use std::sync::Arc;
5 use std::time::{Duration, SystemTime, UNIX_EPOCH};
6
7 use anyhow::{Context, Result, anyhow, bail};
8 use base64::Engine as _;
9 use base64::engine::general_purpose::URL_SAFE_NO_PAD;
10 use oauth2::TokenResponse;
11 use reqwest::Url;
12 use reqwest::header::{CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue};
13 use rmcp::transport::AuthorizationManager;
14 use rmcp::transport::AuthorizationSession;
15 use rmcp::transport::auth::{
16 AuthError, AuthorizationRequest, OAuthClientConfig, OAuthHttpClient, OAuthHttpClientError,
17 OAuthHttpClientFuture, OAuthHttpRedirectPolicy, OAuthHttpRequest, OAuthState,
18 OAuthTokenResponse,
19 };
20 use serde::{Deserialize, Serialize};
21 use sha2::{Digest, Sha256};
22 use tokio::io::{AsyncReadExt, AsyncWriteExt};
23 use tokio::net::TcpListener;
24 use tokio::sync::{Mutex, oneshot};
25 use tokio::time::timeout;
26 use tokio_util::sync::CancellationToken;
27 use urlencoding::decode;
28
29 use super::McpServerConfig;
30
31 const REFRESH_SKEW_MILLIS: u64 = 30_000;
32
33 #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
34 #[serde(rename_all = "snake_case")]
35 pub enum McpAuthStatus {
36 Unsupported,
37 NotLoggedIn,
38 BearerToken,
39 OAuth,
40 }
41
42 impl std::fmt::Display for McpAuthStatus {
43 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
44 let text = match self {
45 Self::Unsupported => "Unsupported",
46 Self::NotLoggedIn => "Not logged in",
47 Self::BearerToken => "Bearer token",
48 Self::OAuth => "OAuth",
49 };
50 f.write_str(text)
51 }
52 }
53
54 /// Context for a failed token refresh. An auth-required failure already
55 /// flips the server to `◆ auth required` and offers the login tool; any
56 /// other failure (a token endpoint answering something the client could
57 /// not parse, a transport error) names the same remedy in words, because
58 /// the operator otherwise sees only the provider's parse error (#5926).
59 /// When the token endpoint did answer, its receipt (status line,
60 /// content-type, masked excerpt) rides along so a provider outage — an HTML
61 /// 502 page — reads differently from a parser defect on JSON it should
62 /// have accepted.
63 fn refresh_failure_context(
64 server_name: &str,
65 names_remedy: bool,
66 receipt: Option<&TokenEndpointReceipt>,
67 ) -> String {
68 if names_remedy {
69 let answered =
70 receipt.map_or_else(String::new, |receipt| format!(" (it answered {receipt})"));
71 format!(
72 "refreshing MCP OAuth token for server {server_name}: the token endpoint did not answer the way the client expects{answered}; \
73 if this persists, run `codewhale mcp login {server_name}` (or `/mcp login {server_name}`) to re-authorize"
74 )
75 } else {
76 format!("refreshing MCP OAuth token for server {server_name}")
77 }
78 }
79
80 /// Longest excerpt of a token-endpoint body a refresh failure keeps.
81 const TOKEN_RECEIPT_EXCERPT_BYTES: usize = 200;
82
83 /// rmcp's own cap on an OAuth response body, mirrored so the recording
84 /// client refuses the same oversized answers the stock one does.
85 const MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES: usize = 1024 * 1024;
86
87 /// Response fields whose values are credentials. Their values are masked
88 /// before any body excerpt is kept; the field names themselves are not
89 /// secrets and stay so the operator can see which fields the answer had.
90 const OAUTH_SECRET_FIELDS: &[&str] = &[
91 "access_token",
92 "refresh_token",
93 "client_secret",
94 "id_token",
95 "authorization",
96 ];
97
98 /// What the token endpoint actually answered. rmcp collapses an
99 /// unparseable answer to `Failed to parse server response` and drops the
100 /// body; this is the receipt it drops, with every credential-shaped value
101 /// masked and the body cut to its first [`TOKEN_RECEIPT_EXCERPT_BYTES`].
102 #[derive(Debug, Clone, PartialEq, Eq)]
103 pub(crate) struct TokenEndpointReceipt {
104 status: u16,
105 reason: Option<&'static str>,
106 content_type: Option<String>,
107 excerpt: String,
108 }
109
110 impl TokenEndpointReceipt {
111 fn from_response(status: reqwest::StatusCode, headers: &HeaderMap, body: &[u8]) -> Self {
112 let content_type = headers
113 .get(CONTENT_TYPE)
114 .and_then(|value| value.to_str().ok())
115 .map(str::trim)
116 .filter(|value| !value.is_empty())
117 .map(str::to_string);
118 Self {
119 status: status.as_u16(),
120 reason: status.canonical_reason(),
121 content_type,
122 excerpt: token_response_excerpt(body),
123 }
124 }
125 }
126
127 impl std::fmt::Display for TokenEndpointReceipt {
128 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
129 write!(f, "HTTP {}", self.status)?;
130 if let Some(reason) = self.reason {
131 write!(f, " {reason}")?;
132 }
133 match self.content_type.as_deref() {
134 Some(content_type) => write!(f, " ({content_type})")?,
135 None => f.write_str(" (no content-type)")?,
136 }
137 if self.excerpt.is_empty() {
138 f.write_str(" with an empty body")
139 } else {
140 write!(f, ": {}", self.excerpt)
141 }
142 }
143 }
144
145 /// Masked, whitespace-collapsed, byte-capped excerpt of a token-endpoint
146 /// body. Masking runs before the cut so a truncated credential never leaks
147 /// its prefix.
148 fn token_response_excerpt(body: &[u8]) -> String {
149 let masked = mask_oauth_secrets(&String::from_utf8_lossy(body));
150 let collapsed = masked.split_whitespace().collect::<Vec<_>>().join(" ");
151 if collapsed.len() <= TOKEN_RECEIPT_EXCERPT_BYTES {
152 return collapsed;
153 }
154 let mut end = TOKEN_RECEIPT_EXCERPT_BYTES;
155 while !collapsed.is_char_boundary(end) {
156 end -= 1;
157 }
158 format!("{}…", &collapsed[..end])
159 }
160
161 fn is_word_byte(byte: u8) -> bool {
162 byte.is_ascii_alphanumeric() || byte == b'_'
163 }
164
165 /// Replace every credential-shaped value in `text` with `***`: JSON
166 /// members (`"access_token": "…"`), form/query pairs (`refresh_token=…`),
167 /// and bearer schemes (`Bearer …`). Field names, separators and everything
168 /// else survive so the shape of the answer stays readable.
169 pub(crate) fn mask_oauth_secrets(text: &str) -> String {
170 let lower = text.to_ascii_lowercase();
171 let bytes = text.as_bytes();
172 let mut out = String::with_capacity(text.len());
173 let mut index = 0;
174 while index < text.len() {
175 let at_word_start = index == 0 || !is_word_byte(bytes[index - 1]);
176 if at_word_start
177 && let Some((value_start, value_end)) = secret_value_span(text, &lower, index)
178 {
179 out.push_str(&text[index..value_start]);
180 out.push_str("***");
181 index = value_end;
182 continue;
183 }
184 let ch = text[index..]
185 .chars()
186 .next()
187 .expect("index sits on a char boundary");
188 out.push(ch);
189 index += ch.len_utf8();
190 }
191 out
192 }
193
194 /// The byte span of the secret value that starts at `start`, if a secret
195 /// field or bearer scheme begins there. Every scan step consumes ASCII
196 /// bytes only, so both ends land on char boundaries.
197 fn secret_value_span(text: &str, lower: &str, start: usize) -> Option<(usize, usize)> {
198 let bytes = text.as_bytes();
199 let skip_spaces = |mut cursor: usize| {
200 while bytes
201 .get(cursor)
202 .is_some_and(|byte| *byte == b' ' || *byte == b'\t')
203 {
204 cursor += 1;
205 }
206 cursor
207 };
208 let unquoted_end = |mut cursor: usize| {
209 while bytes.get(cursor).is_some_and(|byte| {
210 !matches!(byte, b'&' | b',' | b';' | b'}' | b'"' | b'\'') && !byte.is_ascii_whitespace()
211 }) {
212 cursor += 1;
213 }
214 cursor
215 };
216 if lower[start..].starts_with("bearer ") {
217 let value_start = skip_spaces(start + "bearer".len());
218 let value_end = unquoted_end(value_start);
219 return (value_end > value_start).then_some((value_start, value_end));
220 }
221 for field in OAUTH_SECRET_FIELDS {
222 if !lower[start..].starts_with(field) {
223 continue;
224 }
225 let mut cursor = start + field.len();
226 if bytes.get(cursor).is_some_and(|byte| is_word_byte(*byte)) {
227 continue;
228 }
229 if bytes.get(cursor) == Some(&b'"') {
230 cursor += 1;
231 }
232 cursor = skip_spaces(cursor);
233 match bytes.get(cursor) {
234 Some(b':' | b'=') => cursor += 1,
235 _ => continue,
236 }
237 cursor = skip_spaces(cursor);
238 if bytes.get(cursor) == Some(&b'"') {
239 let value_start = cursor + 1;
240 let mut value_end = value_start;
241 while let Some(byte) = bytes.get(value_end) {
242 match byte {
243 b'\\' => value_end += 2,
244 b'"' => break,
245 _ => value_end += 1,
246 }
247 }
248 return Some((value_start, value_end.min(text.len())));
249 }
250 // An unquoted `Authorization: Bearer <token>` carries its scheme in
251 // front of the credential; the whole value is the secret.
252 let value_start = cursor;
253 let mut value_end = unquoted_end(cursor);
254 if matches!(lower[value_start..value_end].as_ref(), "bearer" | "basic")
255 && bytes.get(value_end) == Some(&b' ')
256 {
257 value_end = unquoted_end(skip_spaces(value_end));
258 }
259 return Some((value_start, value_end));
260 }
261 None
262 }
263
264 /// Shared guarded HTTP client for discovery, login and stored credentials.
265 /// It honors each OAuth operation's redirect policy, caps response bodies and keeps
266 /// the receipt of the latest token-endpoint answer (every token request is
267 /// a POST; discovery is GET) so a failed refresh can say what came back.
268 pub(crate) struct RecordingOAuthHttpClient {
269 client: McpHttpClient,
270 last_token_response: std::sync::Mutex<Option<TokenEndpointReceipt>>,
271 }
272
273 impl RecordingOAuthHttpClient {
274 fn new(client: McpHttpClient) -> Self {
275 Self {
276 client,
277 last_token_response: std::sync::Mutex::new(None),
278 }
279 }
280
281 fn take_token_endpoint_receipt(&self) -> Option<TokenEndpointReceipt> {
282 self.last_token_response
283 .lock()
284 .unwrap_or_else(std::sync::PoisonError::into_inner)
285 .take()
286 }
287 }
288
289 impl OAuthHttpClient for RecordingOAuthHttpClient {
290 fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> {
291 Box::pin(async move {
292 let OAuthHttpRequest {
293 request,
294 timeout,
295 redirect_policy,
296 ..
297 } = request;
298 let is_token_request = request.method() == reqwest::Method::POST;
299 let mut request = reqwest::Request::try_from(request)
300 .map_err(|error| Box::new(error) as OAuthHttpClientError)?;
301 if let Some(timeout) = timeout {
302 *request.timeout_mut() = Some(timeout);
303 }
304 let mut response = self
305 .client
306 .execute(
307 request,
308 matches!(redirect_policy, OAuthHttpRedirectPolicy::Follow),
309 )
310 .await
311 .map_err(|error| -> OAuthHttpClientError { error.into() })?;
312 let status = response.status();
313 let version = response.version();
314 let headers = response.headers().clone();
315 let mut body = Vec::new();
316 while let Some(chunk) = response
317 .chunk()
318 .await
319 .map_err(|error| Box::new(error) as OAuthHttpClientError)?
320 {
321 if chunk.len() > MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES - body.len() {
322 return Err(anyhow!(
323 "OAuth HTTP response body exceeds {MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES} bytes"
324 )
325 .into());
326 }
327 body.extend_from_slice(&chunk);
328 }
329 if is_token_request {
330 *self
331 .last_token_response
332 .lock()
333 .unwrap_or_else(std::sync::PoisonError::into_inner) =
334 Some(TokenEndpointReceipt::from_response(status, &headers, &body));
335 }
336 let mut builder = oauth2::http::Response::builder()
337 .status(status)
338 .version(version);
339 for (name, value) in &headers {
340 builder = builder.header(name, value);
341 }
342 builder
343 .body(body)
344 .map_err(|error| Box::new(error) as OAuthHttpClientError)
345 })
346 }
347 }
348
349 pub fn error_looks_auth_required(error: &anyhow::Error) -> bool {
350 error_text_looks_auth_required(&format!("{error:#}"))
351 }
352
353 /// Whether the error chain carries the OAuth `invalid_grant` code: the
354 /// authorization server definitively rejected the presented grant (typically
355 /// a stale or already-rotated refresh token).
356 fn error_is_invalid_grant(error: &anyhow::Error) -> bool {
357 format!("{error:#}")
358 .to_ascii_lowercase()
359 .contains("invalid_grant")
360 }
361
362 /// The one auth-required classifier every surface consults: the pool's
363 /// `◆ auth required` state, the session-boot row, the `/mcp` manager
364 /// recovery verb, and the synthetic `mcp_<server>_authenticate` tool all
365 /// derive from this predicate so a failure is never "needs login" on one
366 /// surface and "failed" on another. `invalid_grant` belongs here because the
367 /// authorization server has definitively rejected the stored grant — only a
368 /// fresh login recovers it.
369 pub fn error_text_looks_auth_required(text: &str) -> bool {
370 let status_401 = text_names_http_status(text, "401");
371 let text = text.to_ascii_lowercase();
372 // `auth required` and `requires oauth` are anchored to the shapes this
373 // product and the Codex-compatible managers actually emit (`◆ auth
374 // required`, `requires OAuth login/authentication/reauthentication`) —
375 // bare substrings would misclassify incidental server errors like
376 // "auth required parameter is missing".
377 status_401
378 || text.contains("unauthorized")
379 || text.contains("authentication_required")
380 || text.contains("invalid_grant")
381 // rmcp 3.2 collapsed several unrecoverable refresh outcomes onto
382 // `AuthError::AuthorizationRequired` (Display: "OAuth authorization
383 // required") — a stored credential with no usable refresh grant, and
384 // every refresh the server definitively rejected. In 2.2 those
385 // arrived as `TokenRefreshFailed("No refresh token available")`, which
386 // no surface recognised. The full phrase is matched so it stays
387 // anchored to rmcp's own wording.
388 || text.contains("oauth authorization required")
389 || text.contains("◆ auth required")
390 || text.contains("requires oauth login")
391 || text.contains("requires oauth authentication")
392 || text.contains("requires oauth reauthentication")
393 || text.contains("not logged in")
394 || text.contains("not-logged-in")
395 || text.contains("re-authorize")
396 || text.contains("/mcp login")
397 || text.contains("mcp login")
398 }
399
400 /// Whether error text names an HTTP status as a status, not as digits inside
401 /// an address or a larger number. Transport errors carry the URL they failed
402 /// on, so a bare substring match read the reset on
403 /// `http://127.0.0.1:50401/mcp` as a 401 and put a healthy server into
404 /// `◆ auth required` — whenever the ephemeral port happened to contain it.
405 pub(crate) fn text_names_http_status(text: &str, status: &str) -> bool {
406 text.split_whitespace()
407 .filter(|token| !token.contains("://"))
408 .any(|token| {
409 token
410 .split(|c: char| !c.is_ascii_digit())
411 .any(|digits| digits == status)
412 })
413 }
414
415 pub fn auth_required_login_hint(server_name: &str) -> String {
416 format!(
417 "MCP server '{server_name}' requires OAuth authentication. Run `codewhale mcp login {server_name}` to authenticate."
418 )
419 }
420
421 /// The one recovery sentence for a server in the `◆ auth required` state,
422 /// chosen by how that server is allowed to authenticate. OAuth-servable
423 /// servers get the login command; plugin-contributed servers (OAuth is
424 /// disabled for them by review policy) and servers with a manual
425 /// Authorization configuration are told which environment-backed
426 /// credential to supply instead, so `/mcp login` is never named for a
427 /// server it would refuse. Environment variable *names* are not secrets;
428 /// their values never appear here.
429 pub(crate) fn auth_required_recovery_hint(server_name: &str, server: &McpServerConfig) -> String {
430 let mut env_vars: Vec<&str> = server
431 .env_headers
432 .values()
433 .map(String::as_str)
434 .chain(server.bearer_token_env_var.as_deref())
435 .collect();
436 env_vars.sort_unstable();
437 env_vars.dedup();
438 let credential_source = if env_vars.is_empty() {
439 "its configured Authorization header".to_string()
440 } else {
441 format!(
442 "the environment variable{} {}",
443 if env_vars.len() == 1 { "" } else { "s" },
444 env_vars.join(", ")
445 )
446 };
447 if let Some(source) = server.reviewed_plugin.as_ref() {
448 return format!(
449 "MCP server '{server_name}' is contributed by plugin '{}' and its credential comes from {credential_source} (OAuth login is disabled for plugin-contributed servers). Set the credential, then run `/mcp reload`.",
450 source.authority.plugin_name
451 );
452 }
453 if server_has_manual_authorization(server) {
454 return format!(
455 "MCP server '{server_name}' authenticates with {credential_source}; the server rejected that credential. Correct it, then run `/mcp reload`."
456 );
457 }
458 auth_required_login_hint(server_name)
459 }
460
461 /// TUI recovery for a stale Streamable HTTP OAuth session. `/mcp auth` is not a
462 /// command; login is `/mcp login <name>` (CLI: `codewhale mcp login <name>`).
463 pub fn tui_reauth_hint() -> &'static str {
464 "Re-authorize this server (/mcp login <name>) to continue."
465 }
466
467 pub fn tui_reauth_refresh_failed_hint() -> &'static str {
468 "Re-authorize this server (/mcp login <name>) or configure a fresh bearer token."
469 }
470
471 #[derive(Debug, Clone, Serialize, Deserialize)]
472 pub struct StoredMcpOAuthTokens {
473 pub server_name: String,
474 pub url: String,
475 pub client_id: String,
476 pub token_response: WrappedOAuthTokenResponse,
477 #[serde(default)]
478 pub expires_at: Option<u64>,
479 }
480
481 impl PartialEq for StoredMcpOAuthTokens {
482 fn eq(&self, other: &Self) -> bool {
483 if self.server_name != other.server_name
484 || self.url != other.url
485 || self.client_id != other.client_id
486 || self.expires_at != other.expires_at
487 {
488 return false;
489 }
490 if self.expires_at.is_none() {
491 return self.token_response == other.token_response;
492 }
493 // Loading a credential derives a decreasing expires_in from the
494 // durable expires_at. That countdown is not a peer token rotation:
495 // comparing it would adopt the same rejected grant after one second
496 // instead of invalidating it. Preserve every other response field.
497 let mut left = self.token_response.clone();
498 let mut right = other.token_response.clone();
499 left.0.set_expires_in(None);
500 right.0.set_expires_in(None);
501 left == right
502 }
503 }
504
505 #[derive(Debug, Clone, Serialize, Deserialize)]
506 pub struct WrappedOAuthTokenResponse(pub OAuthTokenResponse);
507
508 impl PartialEq for WrappedOAuthTokenResponse {
509 fn eq(&self, other: &Self) -> bool {
510 match (serde_json::to_string(self), serde_json::to_string(other)) {
511 (Ok(left), Ok(right)) => left == right,
512 _ => false,
513 }
514 }
515 }
516
517 #[derive(Clone)]
518 pub struct McpOAuthRuntime {
519 inner: Arc<McpOAuthRuntimeInner>,
520 }
521
522 struct McpOAuthRuntimeInner {
523 server_name: String,
524 url: String,
525 manager: Arc<Mutex<AuthorizationManager>>,
526 last_tokens: Mutex<Option<StoredMcpOAuthTokens>>,
527 /// Why the held credential was invalidated (the provider's error code,
528 /// e.g. `invalid_grant`), so every later failure names the cause even
529 /// though the rejected grant is never replayed. `None` while a
530 /// credential is held.
531 rejection: Mutex<Option<String>>,
532 /// The HTTP client the runtime was built with, shared with the manager
533 /// so an adopted on-disk rotation rebuilds it with identical HTTP shape
534 /// and a failed refresh can read the token endpoint's receipt.
535 http_client: Arc<RecordingOAuthHttpClient>,
536 }
537
538 #[derive(Debug, Clone, PartialEq, Eq)]
539 pub struct McpOAuthDiscovery {
540 pub scopes_supported: Option<Vec<String>>,
541 }
542
543 #[derive(Debug, Clone, PartialEq, Eq)]
544 pub struct ResolvedMcpOAuthScopes {
545 pub scopes: Vec<String>,
546 pub source: McpOAuthScopesSource,
547 }
548
549 #[derive(Debug, Clone, Copy, PartialEq, Eq)]
550 pub enum McpOAuthScopesSource {
551 Explicit,
552 Configured,
553 Discovered,
554 Empty,
555 }
556
557 #[derive(Debug, Clone, PartialEq, Eq)]
558 pub struct OAuthProviderError {
559 error: Option<String>,
560 error_description: Option<String>,
561 }
562
563 impl OAuthProviderError {
564 fn new(error: Option<String>, error_description: Option<String>) -> Self {
565 Self {
566 error,
567 error_description,
568 }
569 }
570 }
571
572 impl std::fmt::Display for OAuthProviderError {
573 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
574 match (self.error.as_deref(), self.error_description.as_deref()) {
575 (Some(error), Some(description)) => {
576 write!(f, "OAuth provider returned `{error}`: {description}")
577 }
578 (Some(error), None) => write!(f, "OAuth provider returned `{error}`"),
579 (None, Some(description)) => write!(f, "OAuth error: {description}"),
580 (None, None) => write!(f, "OAuth provider returned an error"),
581 }
582 }
583 }
584
585 impl std::error::Error for OAuthProviderError {}
586
587 /// Build an `AuthorizationManager` preloaded with stored credentials, the
588 /// shared construction step for initial load and for adopting a credential
589 /// that another process rotated on disk.
590 async fn manager_from_stored_tokens(
591 url: &str,
592 tokens: &StoredMcpOAuthTokens,
593 http_client: &Arc<RecordingOAuthHttpClient>,
594 ) -> Result<AuthorizationManager> {
595 let client = Arc::clone(http_client) as Arc<dyn OAuthHttpClient>;
596 let mut state = OAuthState::new_with_oauth_http_client(url.to_string(), client).await?;
597 state
598 .set_credentials(&tokens.client_id, tokens.token_response.0.clone())
599 .await
600 .context("installing stored MCP OAuth credentials")?;
601
602 match state {
603 OAuthState::Authorized(manager) | OAuthState::Unauthorized(manager) => Ok(manager),
604 _ => bail!("unexpected MCP OAuth state while preparing stored credentials"),
605 }
606 }
607
608 impl McpOAuthRuntime {
609 #[cfg(test)]
610 pub(super) async fn from_server_config(
611 server_name: &str,
612 server: &McpServerConfig,
613 default_headers: HeaderMap,
614 ) -> Result<Option<Self>> {
615 if server.reviewed_plugin.is_some() || server_has_manual_authorization(server) {
616 return Ok(None);
617 }
618 let Some(url) = server.url.as_deref() else {
619 return Ok(None);
620 };
621 let client = oauth_http_client(server, url, None)?;
622 Self::from_server_config_with_client(server_name, server, default_headers, client).await
623 }
624
625 pub(super) async fn from_server_config_with_client(
626 server_name: &str,
627 server: &McpServerConfig,
628 default_headers: HeaderMap,
629 client: McpHttpClient,
630 ) -> Result<Option<Self>> {
631 if server.reviewed_plugin.is_some() {
632 return Ok(None);
633 }
634 let Some(url) = server.url.as_deref() else {
635 return Ok(None);
636 };
637 if server_has_manual_authorization(server) {
638 return Ok(None);
639 }
640 let Some(tokens) = load_oauth_tokens(server_name, url)? else {
641 return Ok(None);
642 };
643 Self::from_stored_tokens(
644 server_name,
645 url,
646 tokens,
647 client.with_default_headers(default_headers),
648 )
649 .await
650 .map(Some)
651 }
652
653 async fn from_stored_tokens(
654 server_name: &str,
655 url: &str,
656 mut tokens: StoredMcpOAuthTokens,
657 client: McpHttpClient,
658 ) -> Result<Self> {
659 refresh_expires_in_from_timestamp(&mut tokens);
660 let client = client.with_credential_transport()?;
661 let http_client = Arc::new(RecordingOAuthHttpClient::new(client));
662 let manager = manager_from_stored_tokens(url, &tokens, &http_client).await?;
663
664 Ok(Self {
665 inner: Arc::new(McpOAuthRuntimeInner {
666 server_name: server_name.to_string(),
667 url: url.to_string(),
668 manager: Arc::new(Mutex::new(manager)),
669 last_tokens: Mutex::new(Some(tokens)),
670 rejection: Mutex::new(None),
671 http_client,
672 }),
673 })
674 }
675
676 pub async fn authorization_header(&self) -> Result<Option<String>> {
677 self.refresh_if_needed().await?;
678 // Never send a credential the provider already rejected; the request
679 // goes out unauthenticated and the server's 401 drives the reactive
680 // refresh, which adopts a peer's login or reports auth-required.
681 if self.is_invalidated().await {
682 return Ok(None);
683 }
684 let credentials = {
685 let guard = self.inner.manager.lock().await;
686 let (_client_id, credentials) = guard
687 .get_credentials()
688 .await
689 .context("reading MCP OAuth credentials")?;
690 credentials
691 };
692 let Some(credentials) = credentials else {
693 return Ok(None);
694 };
695 let token = credentials.access_token().secret().trim();
696 if token.is_empty() {
697 Ok(None)
698 } else {
699 Ok(Some(format!("Bearer {token}")))
700 }
701 }
702
703 async fn refresh_if_needed(&self) -> Result<()> {
704 let expires_at = {
705 let guard = self.inner.last_tokens.lock().await;
706 guard.as_ref().and_then(|tokens| tokens.expires_at)
707 };
708 if !token_needs_refresh(expires_at) {
709 return Ok(());
710 }
711 self.refresh_and_persist().await
712 }
713
714 /// Force a token refresh regardless of the local expiry clock (T4): a
715 /// 401/403 means the server no longer accepts the token — clock skew,
716 /// server-side revocation, or rotation — so the expiry-based gate must
717 /// not decide alone.
718 pub(crate) async fn force_refresh(&self) -> Result<()> {
719 self.refresh_and_persist().await
720 }
721
722 /// Whether this runtime's credential was definitively rejected by the
723 /// provider and invalidated. `last_tokens` is `Some` from construction
724 /// and after every persisted refresh; only [`Self::clear_stored_tokens`]
725 /// empties it.
726 async fn is_invalidated(&self) -> bool {
727 self.inner.last_tokens.lock().await.is_none()
728 }
729
730 async fn refresh_and_persist(&self) -> Result<()> {
731 // A credential the provider definitively rejected is never replayed:
732 // the `AuthorizationManager` still holds it, but every later refresh
733 // with that grant is a guaranteed `invalid_grant`. The only way back
734 // is a credential another process stored since (a completed login),
735 // so adopt that when present and otherwise report auth-required
736 // without touching the token endpoint.
737 if self.is_invalidated().await {
738 if !self.adopt_rotated_on_disk_tokens().await? {
739 let reason = self
740 .inner
741 .rejection
742 .lock()
743 .await
744 .clone()
745 .unwrap_or_else(|| "unauthorized".to_string());
746 bail!(
747 "stored MCP OAuth credential for server {} was rejected by the provider ({reason}) and removed; the server requires OAuth login again",
748 self.inner.server_name
749 );
750 }
751 let adopted_needs_refresh = {
752 let last = self.inner.last_tokens.lock().await;
753 token_needs_refresh(last.as_ref().and_then(|tokens| tokens.expires_at))
754 };
755 if !adopted_needs_refresh {
756 return Ok(());
757 }
758 }
759 // Only this refresh's answer may explain this refresh's failure.
760 self.inner.http_client.take_token_endpoint_receipt();
761 let mut err = match self.try_refresh_and_persist().await {
762 Ok(()) => return Ok(()),
763 Err(err) => err,
764 };
765 // Refresh-race tolerance: another codewhale process sharing this token
766 // store (a concurrent `mcp login`, or a peer session's refresh) may
767 // have rotated the credential after this runtime loaded its copy, and
768 // single-use refresh tokens then fail here with `invalid_grant`.
769 // Re-read the store once; when the on-disk credential changed, adopt
770 // it — using it directly while fresh, or retrying the refresh exactly
771 // once with the rotated grant — before surfacing failure. When the
772 // store is unchanged the grant is simply dead: the auth-required
773 // branch below invalidates it so the server flips to `◆ auth
774 // required` and the self-serve login tool appears, instead of every
775 // later connect replaying the same rejected refresh.
776 if error_is_invalid_grant(&err) && self.adopt_rotated_on_disk_tokens().await? {
777 let adopted_needs_refresh = {
778 let last = self.inner.last_tokens.lock().await;
779 token_needs_refresh(last.as_ref().and_then(|tokens| tokens.expires_at))
780 };
781 if !adopted_needs_refresh {
782 return Ok(());
783 }
784 match self.try_refresh_and_persist().await {
785 Ok(()) => return Ok(()),
786 Err(retry_err) => err = retry_err,
787 }
788 }
789 if error_looks_auth_required(&err) {
790 let reason = if error_is_invalid_grant(&err) {
791 "invalid_grant"
792 } else {
793 "unauthorized"
794 };
795 self.clear_stored_tokens(reason).await?;
796 }
797 let server_name = self.inner.server_name.clone();
798 let names_remedy = !error_looks_auth_required(&err);
799 let receipt = self.inner.http_client.take_token_endpoint_receipt();
800 Err(err)
801 .with_context(|| refresh_failure_context(&server_name, names_remedy, receipt.as_ref()))
802 }
803
804 async fn try_refresh_and_persist(&self) -> Result<()> {
805 let refresh_result = {
806 let guard = self.inner.manager.lock().await;
807 guard.refresh_token().await
808 };
809 refresh_result.map_err(|err| anyhow!(err))?;
810 self.persist_if_needed().await
811 }
812
813 /// Re-read the on-disk credential after an `invalid_grant` refresh
814 /// failure and, when another process rotated it, rebuild the manager
815 /// around the rotated token exactly like initial construction. Returns
816 /// `true` only when the stored credential actually changed; an unchanged
817 /// store means the failure is ours to report.
818 async fn adopt_rotated_on_disk_tokens(&self) -> Result<bool> {
819 let Some(stored) = load_oauth_tokens(&self.inner.server_name, &self.inner.url)? else {
820 return Ok(false);
821 };
822 let changed = {
823 let last = self.inner.last_tokens.lock().await;
824 last.as_ref() != Some(&stored)
825 };
826 if !changed {
827 return Ok(false);
828 }
829 let manager =
830 manager_from_stored_tokens(&self.inner.url, &stored, &self.inner.http_client).await?;
831 *self.inner.manager.lock().await = manager;
832 *self.inner.last_tokens.lock().await = Some(stored);
833 *self.inner.rejection.lock().await = None;
834 Ok(true)
835 }
836
837 /// Invalidate the credential this runtime holds after the provider
838 /// definitively rejected it. Never deletes a newer durable winner: when
839 /// the on-disk credential no longer matches the one we hold, another
840 /// process rotated it after our copy loaded, and that credential — not
841 /// ours — is the one the next connect must try. The manager is rebuilt
842 /// around that winner exactly like initial construction; remembering
843 /// the rotated token while keeping our dead grant would make the next
844 /// `invalid_grant` compare an "unchanged" store and delete the newer
845 /// valid credential.
846 ///
847 /// The comparison and the delete are one step under the secret store's
848 /// entry lock (the same lock every save takes), so a login that lands
849 /// between them is never the credential that gets deleted.
850 async fn clear_stored_tokens(&self, reason: &str) -> Result<()> {
851 let held = { self.inner.last_tokens.lock().await.take() };
852 let Some(held) = held else {
853 return Ok(());
854 };
855 match delete_oauth_tokens_if_held(&self.inner.server_name, &self.inner.url, &held)? {
856 Some(stored) => {
857 tracing::debug!(
858 target: "mcp",
859 server = %self.inner.server_name,
860 "MCP OAuth credential was rotated by another process; keeping the on-disk winner"
861 );
862 let manager =
863 manager_from_stored_tokens(&self.inner.url, &stored, &self.inner.http_client)
864 .await?;
865 *self.inner.manager.lock().await = manager;
866 *self.inner.last_tokens.lock().await = Some(stored);
867 *self.inner.rejection.lock().await = None;
868 }
869 None => {
870 *self.inner.rejection.lock().await = Some(reason.to_string());
871 }
872 }
873 Ok(())
874 }
875
876 async fn persist_if_needed(&self) -> Result<()> {
877 let (client_id, credentials) = {
878 let guard = self.inner.manager.lock().await;
879 guard
880 .get_credentials()
881 .await
882 .context("reading refreshed MCP OAuth credentials")?
883 };
884 let Some(credentials) = credentials else {
885 let mut last = self.inner.last_tokens.lock().await;
886 if let Some(previous) = last.take() {
887 // Only our own credential goes; a newer login stays.
888 delete_oauth_tokens_if_held(&self.inner.server_name, &self.inner.url, &previous)?;
889 }
890 return Ok(());
891 };
892
893 let new_response = WrappedOAuthTokenResponse(credentials.clone());
894 let mut last = self.inner.last_tokens.lock().await;
895 let same_token = last
896 .as_ref()
897 .map(|previous| previous.token_response == new_response)
898 .unwrap_or(false);
899 let expires_at = if same_token {
900 last.as_ref().and_then(|previous| previous.expires_at)
901 } else {
902 compute_expires_at_millis(&credentials)
903 };
904 let stored = StoredMcpOAuthTokens {
905 server_name: self.inner.server_name.clone(),
906 url: self.inner.url.clone(),
907 client_id,
908 token_response: new_response,
909 expires_at,
910 };
911 if last.as_ref() != Some(&stored) {
912 save_oauth_tokens(&stored)?;
913 *last = Some(stored);
914 }
915 Ok(())
916 }
917 }
918
919 pub async fn auth_status_for_server(
920 name: &str,
921 server: &McpServerConfig,
922 network_policy: Option<&NetworkPolicyDecider>,
923 ) -> McpAuthStatus {
924 if server.reviewed_plugin.is_some() || !server.is_enabled() || server.url.is_none() {
925 return McpAuthStatus::Unsupported;
926 }
927 if server_has_manual_authorization(server) {
928 return McpAuthStatus::BearerToken;
929 }
930 let Some(url) = server.url.as_deref() else {
931 return McpAuthStatus::Unsupported;
932 };
933 match load_oauth_tokens(name, url) {
934 Ok(Some(tokens)) if oauth_tokens_are_usable(&tokens) => return McpAuthStatus::OAuth,
935 Ok(Some(_)) => return McpAuthStatus::NotLoggedIn,
936 Ok(None) => {}
937 Err(err) => {
938 tracing::warn!(target: "mcp", server = %name, error = %err, "failed to read MCP OAuth tokens");
939 }
940 }
941
942 let headers = match build_default_headers(&server.headers, &server.env_headers) {
943 Ok(headers) => headers,
944 Err(err) => {
945 tracing::warn!(target: "mcp", server = %name, error = %err, "failed to build MCP OAuth discovery headers");
946 return McpAuthStatus::Unsupported;
947 }
948 };
949 match discover_streamable_http_oauth_for_server(server, url, headers, network_policy).await {
950 Ok(Some(_)) => McpAuthStatus::NotLoggedIn,
951 Ok(None) => McpAuthStatus::Unsupported,
952 Err(err) => {
953 tracing::debug!(target: "mcp", server = %name, error = %err, "MCP OAuth discovery failed");
954 McpAuthStatus::Unsupported
955 }
956 }
957 }
958
959 pub async fn oauth_login_support(
960 server: &McpServerConfig,
961 network_policy: Option<&NetworkPolicyDecider>,
962 ) -> Result<Option<McpOAuthDiscovery>> {
963 if server.reviewed_plugin.is_some() {
964 return Ok(None);
965 }
966 let Some(url) = server.url.as_deref() else {
967 return Ok(None);
968 };
969 if server_has_manual_authorization(server) {
970 return Ok(None);
971 }
972 let headers = build_default_headers(&server.headers, &server.env_headers)?;
973 discover_streamable_http_oauth_for_server(server, url, headers, network_policy).await
974 }
975
976 fn oauth_http_client(
977 server: &McpServerConfig,
978 url: &str,
979 network_policy: Option<&NetworkPolicyDecider>,
980 ) -> Result<McpHttpClient> {
981 let timeouts = super::McpTimeouts::default();
982 McpHttpClient::new(
983 url,
984 server.runtime_added,
985 server.reviewed_plugin.is_some(),
986 server.allow_private_network,
987 network_policy,
988 Duration::from_secs(server.effective_connect_timeout(&timeouts)),
989 Duration::from_secs(server.effective_read_timeout(&timeouts)),
990 )
991 .and_then(McpHttpClient::with_credential_transport)
992 }
993
994 fn oauth_login_client(
995 server: &McpServerConfig,
996 url: &str,
997 network_policy: Option<&NetworkPolicyDecider>,
998 ) -> Result<McpHttpClient> {
999 let headers = build_default_headers(&server.headers, &server.env_headers)?;
1000 Ok(oauth_http_client(server, url, network_policy)?.with_default_headers(headers))
1001 }
1002
1003 async fn discover_streamable_http_oauth_for_server(
1004 server: &McpServerConfig,
1005 url: &str,
1006 default_headers: HeaderMap,
1007 network_policy: Option<&NetworkPolicyDecider>,
1008 ) -> Result<Option<McpOAuthDiscovery>> {
1009 let client =
1010 oauth_http_client(server, url, network_policy)?.with_default_headers(default_headers);
1011 discover_streamable_http_oauth_with_client(url, client).await
1012 }
1013
1014 async fn discover_streamable_http_oauth_with_client(
1015 url: &str,
1016 client: McpHttpClient,
1017 ) -> Result<Option<McpOAuthDiscovery>> {
1018 let client = Arc::new(RecordingOAuthHttpClient::new(client));
1019 let manager = AuthorizationManager::new_with_oauth_http_client(url, client).await?;
1020 match tokio::time::timeout(Duration::from_secs(5), manager.resolve_metadata()).await? {
1021 Ok(resolution) => Ok(Some(McpOAuthDiscovery {
1022 scopes_supported: normalize_scopes(resolution.metadata.scopes_supported),
1023 })),
1024 Err(AuthError::NoAuthorizationSupport) => Ok(None),
1025 Err(err) => Err(err.into()),
1026 }
1027 }
1028
1029 pub fn resolve_oauth_scopes(
1030 explicit_scopes: Option<Vec<String>>,
1031 configured_scopes: Vec<String>,
1032 discovered_scopes: Option<Vec<String>>,
1033 ) -> ResolvedMcpOAuthScopes {
1034 if let Some(scopes) = explicit_scopes {
1035 return ResolvedMcpOAuthScopes {
1036 scopes,
1037 source: McpOAuthScopesSource::Explicit,
1038 };
1039 }
1040 if !configured_scopes.is_empty() {
1041 return ResolvedMcpOAuthScopes {
1042 scopes: configured_scopes,
1043 source: McpOAuthScopesSource::Configured,
1044 };
1045 }
1046 if let Some(scopes) = discovered_scopes
1047 && !scopes.is_empty()
1048 {
1049 return ResolvedMcpOAuthScopes {
1050 scopes,
1051 source: McpOAuthScopesSource::Discovered,
1052 };
1053 }
1054 ResolvedMcpOAuthScopes {
1055 scopes: Vec::new(),
1056 source: McpOAuthScopesSource::Empty,
1057 }
1058 }
1059
1060 pub async fn perform_oauth_login_for_server(
1061 name: &str,
1062 server: &McpServerConfig,
1063 explicit_scopes: Option<Vec<String>>,
1064 callback_port: Option<u16>,
1065 callback_url: Option<&str>,
1066 network_policy: Option<&NetworkPolicyDecider>,
1067 ) -> Result<()> {
1068 perform_oauth_login_for_server_with_cancel(
1069 name,
1070 server,
1071 explicit_scopes,
1072 callback_port,
1073 callback_url,
1074 CancellationToken::new(),
1075 network_policy,
1076 )
1077 .await
1078 }
1079
1080 /// Run an MCP OAuth login that can be stopped by the caller.
1081 ///
1082 /// Cancellation drops the in-flight OAuth future before this function returns,
1083 /// which also closes its callback listener. A caller that replaces one login
1084 /// with another should await the cancelled call before starting the replacement.
1085 pub async fn perform_oauth_login_for_server_with_cancel(
1086 name: &str,
1087 server: &McpServerConfig,
1088 explicit_scopes: Option<Vec<String>>,
1089 callback_port: Option<u16>,
1090 callback_url: Option<&str>,
1091 cancellation_token: CancellationToken,
1092 network_policy: Option<&NetworkPolicyDecider>,
1093 ) -> Result<()> {
1094 if server.reviewed_plugin.is_some() {
1095 bail!(
1096 "OAuth is disabled for plugin-contributed MCP servers; use a reviewed environment-backed header or bearer token"
1097 );
1098 }
1099 run_cancellable_oauth(
1100 &cancellation_token,
1101 perform_oauth_login_for_server_inner(
1102 name,
1103 server,
1104 explicit_scopes,
1105 callback_port,
1106 callback_url,
1107 network_policy,
1108 ),
1109 )
1110 .await
1111 }
1112
1113 async fn run_cancellable_oauth<F, T>(cancellation_token: &CancellationToken, future: F) -> Result<T>
1114 where
1115 F: std::future::Future<Output = Result<T>>,
1116 {
1117 tokio::select! {
1118 biased;
1119 _ = cancellation_token.cancelled() => bail!("OAuth login was cancelled"),
1120 result = future => result,
1121 }
1122 }
1123
1124 /// Shared gate + scope resolution for `/mcp login` and the model-driven
1125 /// authenticate tool: URL-based servers only, no manual Authorization config,
1126 /// scopes from explicit argument, config, or discovery (in that order).
1127 async fn resolve_oauth_login(
1128 name: &str,
1129 server: &McpServerConfig,
1130 explicit_scopes: Option<Vec<String>>,
1131 network_policy: Option<&NetworkPolicyDecider>,
1132 ) -> Result<(String, ResolvedMcpOAuthScopes)> {
1133 let Some(url) = server.url.as_deref() else {
1134 bail!("OAuth login is only supported for URL-based MCP servers");
1135 };
1136 if server_has_manual_authorization(server) {
1137 bail!("MCP server '{name}' already has bearer/static Authorization configured");
1138 }
1139
1140 let discovery = if explicit_scopes.is_none() && server.scopes.is_empty() {
1141 oauth_login_support(server, network_policy).await?
1142 } else {
1143 None
1144 };
1145 let resolved_scopes = resolve_oauth_scopes(
1146 explicit_scopes,
1147 server.scopes.clone(),
1148 discovery.and_then(|discovery| discovery.scopes_supported),
1149 );
1150 Ok((url.to_string(), resolved_scopes))
1151 }
1152
1153 async fn perform_oauth_login_for_server_inner(
1154 name: &str,
1155 server: &McpServerConfig,
1156 explicit_scopes: Option<Vec<String>>,
1157 callback_port: Option<u16>,
1158 callback_url: Option<&str>,
1159 network_policy: Option<&NetworkPolicyDecider>,
1160 ) -> Result<()> {
1161 let (url, resolved_scopes) =
1162 resolve_oauth_login(name, server, explicit_scopes, network_policy).await?;
1163
1164 match perform_oauth_login(
1165 name,
1166 &url,
1167 oauth_login_client(server, &url, network_policy)?,
1168 &resolved_scopes.scopes,
1169 server.oauth_client_id(),
1170 server.oauth_resource.as_deref(),
1171 callback_port,
1172 callback_url,
1173 )
1174 .await
1175 {
1176 Ok(()) => Ok(()),
1177 Err(err)
1178 if resolved_scopes.source == McpOAuthScopesSource::Discovered
1179 && err.downcast_ref::<OAuthProviderError>().is_some() =>
1180 {
1181 println!("OAuth provider rejected discovered scopes. Retrying without scopes...");
1182 perform_oauth_login(
1183 name,
1184 &url,
1185 oauth_login_client(server, &url, network_policy)?,
1186 &[],
1187 server.oauth_client_id(),
1188 server.oauth_resource.as_deref(),
1189 callback_port,
1190 callback_url,
1191 )
1192 .await
1193 }
1194 Err(err) => Err(err),
1195 }
1196 }
1197
1198 #[allow(clippy::too_many_arguments)]
1199 async fn perform_oauth_login(
1200 server_name: &str,
1201 server_url: &str,
1202 client: McpHttpClient,
1203 scopes: &[String],
1204 oauth_client_id: Option<&str>,
1205 oauth_resource: Option<&str>,
1206 callback_port: Option<u16>,
1207 callback_url: Option<&str>,
1208 ) -> Result<()> {
1209 OauthLoginFlow::new(
1210 server_name,
1211 server_url,
1212 client,
1213 scopes,
1214 oauth_client_id,
1215 oauth_resource,
1216 callback_port,
1217 callback_url,
1218 )
1219 .await?
1220 .finish()
1221 .await
1222 }
1223
1224 /// How an OAuth login announces its authorization URL.
1225 #[derive(Debug, Clone, Copy, PartialEq, Eq)]
1226 enum OAuthLoginAnnounce {
1227 /// `/mcp login` in a terminal: print the URL and open a browser.
1228 Terminal,
1229 /// Model-driven `mcp_<server>_authenticate` tool: never write to stdout
1230 /// (a tool call inside a running session is not a terminal); the result
1231 /// carries the URL for the model to relay to the user verbatim.
1232 Tool { open_browser: bool },
1233 }
1234
1235 /// An in-flight OAuth login started by the model-driven authenticate tool.
1236 ///
1237 /// The authorization URL is available immediately so the model can relay it
1238 /// to the user verbatim; [`McpOAuthToolLogin::finish`] then blocks on the
1239 /// loopback callback (up to 5 minutes, same as `/mcp login`) and persists
1240 /// the issued tokens to the shared store on success.
1241 pub struct McpOAuthToolLogin {
1242 server_name: String,
1243 server: McpServerConfig,
1244 scopes_source: McpOAuthScopesSource,
1245 network_policy: Option<NetworkPolicyDecider>,
1246 flow: OauthLoginFlow,
1247 open_browser: bool,
1248 }
1249
1250 impl McpOAuthToolLogin {
1251 /// The exact authorization URL the user must visit. Treat it as
1252 /// sensitive: never modify it or strip query parameters.
1253 #[must_use]
1254 pub fn authorization_url(&self) -> &str {
1255 &self.flow.auth_url
1256 }
1257
1258 /// Block on the browser callback and persist the issued tokens. Mirrors
1259 /// the `/mcp login` retry: when the provider rejects scopes that came
1260 /// from discovery (rather than explicit config), restart once without
1261 /// scopes.
1262 pub async fn finish(self) -> Result<()> {
1263 let announce = OAuthLoginAnnounce::Tool {
1264 open_browser: self.open_browser,
1265 };
1266 let retry_without_scopes = self.scopes_source == McpOAuthScopesSource::Discovered;
1267 match self.flow.finish_with_announce(announce).await {
1268 Ok(()) => Ok(()),
1269 Err(err)
1270 if retry_without_scopes && err.downcast_ref::<OAuthProviderError>().is_some() =>
1271 {
1272 let server = &self.server;
1273 let url = server
1274 .url
1275 .as_deref()
1276 .expect("tool login is gated to URL-based servers at begin");
1277 OauthLoginFlow::new(
1278 &self.server_name,
1279 url,
1280 oauth_login_client(server, url, self.network_policy.as_ref())?,
1281 &[],
1282 server.oauth_client_id(),
1283 server.oauth_resource.as_deref(),
1284 None,
1285 None,
1286 )
1287 .await?
1288 .finish_with_announce(announce)
1289 .await
1290 }
1291 Err(err) => Err(err),
1292 }
1293 }
1294 }
1295
1296 /// Begin the same OAuth login flow `/mcp login` runs, for the model-driven
1297 /// `mcp_<server>_authenticate` tool. The callback listener binds an
1298 /// ephemeral loopback port; callers needing a pre-registered redirect URI
1299 /// keep the terminal `/mcp login <name>` path, which honors the configured
1300 /// callback overrides.
1301 pub async fn begin_oauth_login_for_server_tool(
1302 name: &str,
1303 server: &McpServerConfig,
1304 explicit_scopes: Option<Vec<String>>,
1305 callback_port: Option<u16>,
1306 callback_url: Option<&str>,
1307 network_policy: Option<&NetworkPolicyDecider>,
1308 ) -> Result<McpOAuthToolLogin> {
1309 if server.reviewed_plugin.is_some() {
1310 bail!(
1311 "OAuth is disabled for plugin-contributed MCP servers; use a reviewed environment-backed header or bearer token"
1312 );
1313 }
1314 let (url, resolved_scopes) =
1315 resolve_oauth_login(name, server, explicit_scopes, network_policy).await?;
1316 let flow = OauthLoginFlow::new(
1317 name,
1318 &url,
1319 oauth_login_client(server, &url, network_policy)?,
1320 &resolved_scopes.scopes,
1321 server.oauth_client_id(),
1322 server.oauth_resource.as_deref(),
1323 callback_port,
1324 callback_url,
1325 )
1326 .await?;
1327 Ok(McpOAuthToolLogin {
1328 server_name: name.to_string(),
1329 server: server.clone(),
1330 scopes_source: resolved_scopes.source,
1331 network_policy: network_policy.cloned(),
1332 flow,
1333 // The test build drives the loopback callback itself; a real browser
1334 // launch from a unit test would hijack the developer's desktop.
1335 open_browser: !cfg!(test),
1336 })
1337 }
1338
1339 /// Whether the self-serve OAuth login flow can run for this server at all:
1340 /// URL-based, not plugin-contributed (plugin servers authenticate through
1341 /// reviewed environment-backed headers, and OAuth storage is disabled for
1342 /// them), and without a manual Authorization configuration that an OAuth
1343 /// login would conflict with.
1344 pub(crate) fn server_supports_oauth_login(server: &McpServerConfig) -> bool {
1345 server.reviewed_plugin.is_none()
1346 && server.url.is_some()
1347 && !server_has_manual_authorization(server)
1348 }
1349
1350 /// Whether the shared token store already holds a usable credential for this
1351 /// server — i.e. a login completed in another process since the caller last
1352 /// checked. Used by the authenticate tool's already-authorized branch.
1353 pub(crate) fn has_usable_stored_tokens(name: &str, server: &McpServerConfig) -> bool {
1354 let Some(url) = server.url.as_deref() else {
1355 return false;
1356 };
1357 load_oauth_tokens(name, url)
1358 .ok()
1359 .flatten()
1360 .is_some_and(|tokens| oauth_tokens_are_usable(&tokens))
1361 }
1362
1363 /// Model-facing description for the synthetic `mcp_<server>_authenticate`
1364 /// tool. The coaching contract (show the URL verbatim, the call blocks, real
1365 /// tools replace this one on success) is pinned by tests.
1366 pub(crate) fn authenticate_tool_description(server_name: &str) -> String {
1367 format!(
1368 "Authenticate with MCP server \"{server_name}\" via OAuth.\n\n\
1369 This server requires an OAuth login that has not yet been completed, so its \
1370 real tools are currently unavailable. Calling this tool starts the \
1371 authorization flow:\n\n\
1372 1. A browser window is opened for the user to sign in and approve the \
1373 Codewhale client, and the exact authorization URL is shown to the user in \
1374 the session status while this call waits. The same URL is returned in this \
1375 call's result; if the user reports the browser did not open, show that URL \
1376 to the user verbatim and ask them to complete the sign-in there.\n\
1377 2. The call blocks (up to 5 minutes) until the browser flow completes on the \
1378 local callback listener, is declined, or times out. Do not assume success \
1379 before the call returns.\n\
1380 3. On success the server reconnects and its real MCP tools replace this \
1381 synthetic authenticate tool, becoming callable from the next model request \
1382 in this session.\n\n\
1383 Treat the URL as sensitive — do not modify it or strip query parameters. If \
1384 the flow is declined, cancelled, or times out, the call returns an error; \
1385 relay it truthfully and suggest `/mcp login {server_name}` in the TUI or \
1386 `codewhale mcp login {server_name}` from a terminal."
1387 )
1388 }
1389
1390 pub fn delete_oauth_tokens_for_server(name: &str, server: &McpServerConfig) -> Result<bool> {
1391 if server.reviewed_plugin.is_some() {
1392 bail!("OAuth storage is disabled for plugin-contributed MCP servers");
1393 }
1394 let Some(url) = server.url.as_deref() else {
1395 bail!("OAuth logout is only supported for URL-based MCP servers");
1396 };
1397 delete_oauth_tokens(name, url)
1398 }
1399
1400 pub(crate) fn server_has_manual_authorization(server: &McpServerConfig) -> bool {
1401 server.bearer_token_env_var.is_some()
1402 || contains_authorization_header(&server.headers)
1403 || contains_authorization_header(&server.env_headers)
1404 }
1405
1406 pub fn build_default_headers(
1407 http_headers: &HashMap<String, String>,
1408 env_headers: &HashMap<String, String>,
1409 ) -> Result<HeaderMap> {
1410 let mut headers = HeaderMap::new();
1411 for (name, value) in http_headers {
1412 insert_header(&mut headers, name, value)?;
1413 }
1414 for (name, env_var) in env_headers {
1415 if let Ok(value) = std::env::var(env_var)
1416 && !value.trim().is_empty()
1417 {
1418 insert_header(&mut headers, name, &value)?;
1419 }
1420 }
1421 Ok(headers)
1422 }
1423
1424 fn insert_header(headers: &mut HeaderMap, name: &str, value: &str) -> Result<()> {
1425 if !super::headers::is_safe_custom_header(name, value) {
1426 bail!("unsafe MCP HTTP header '{name}'");
1427 }
1428 let name = HeaderName::from_bytes(name.as_bytes())
1429 .with_context(|| format!("invalid MCP HTTP header name '{name}'"))?;
1430 let value = HeaderValue::from_str(value).with_context(|| "invalid MCP HTTP header value")?;
1431 headers.insert(name, value);
1432 Ok(())
1433 }
1434
1435 fn contains_authorization_header(headers: &HashMap<String, String>) -> bool {
1436 headers
1437 .keys()
1438 .any(|key| key.trim().eq_ignore_ascii_case("authorization"))
1439 }
1440
1441 fn normalize_scopes(scopes_supported: Option<Vec<String>>) -> Option<Vec<String>> {
1442 let scopes_supported = scopes_supported?;
1443 let mut normalized = Vec::new();
1444 for scope in scopes_supported {
1445 let scope = scope.trim();
1446 if scope.is_empty() {
1447 continue;
1448 }
1449 let scope = scope.to_string();
1450 if !normalized.contains(&scope) {
1451 normalized.push(scope);
1452 }
1453 }
1454 (!normalized.is_empty()).then_some(normalized)
1455 }
1456
1457 pub(crate) fn load_oauth_tokens(
1458 server_name: &str,
1459 url: &str,
1460 ) -> Result<Option<StoredMcpOAuthTokens>> {
1461 let secrets = codewhale_secrets::Secrets::auto_detect();
1462 let key = store_key(server_name, url);
1463 let Some(serialized) = secrets
1464 .get(&key)
1465 .with_context(|| format!("reading MCP OAuth token for '{server_name}'"))?
1466 else {
1467 return Ok(None);
1468 };
1469 decode_stored_oauth_tokens(&serialized, server_name).map(Some)
1470 }
1471
1472 fn decode_stored_oauth_tokens(serialized: &str, server_name: &str) -> Result<StoredMcpOAuthTokens> {
1473 let mut tokens = parse_stored_oauth_tokens(serialized, server_name)?;
1474 refresh_expires_in_from_timestamp(&mut tokens);
1475 Ok(tokens)
1476 }
1477
1478 /// Delete the stored credential only while it is still exactly `held`,
1479 /// comparing and deleting under the store's entry lock. Returns the different
1480 /// credential found in its place (left untouched), or `None` when the entry
1481 /// was `held` and is now gone, or was already absent. An unreadable entry is
1482 /// an error and is left in place, as a plain load would report it.
1483 fn delete_oauth_tokens_if_held(
1484 server_name: &str,
1485 url: &str,
1486 held: &StoredMcpOAuthTokens,
1487 ) -> Result<Option<StoredMcpOAuthTokens>> {
1488 let secrets = codewhale_secrets::Secrets::auto_detect();
1489 let key = store_key(server_name, url);
1490 secrets
1491 .with_entry_transaction(&key, |current| {
1492 let Some(serialized) = current.as_deref() else {
1493 return Ok(Ok(None));
1494 };
1495 Ok(match decode_stored_oauth_tokens(serialized, server_name) {
1496 Ok(stored) if stored == *held => {
1497 *current = None;
1498 Ok(None)
1499 }
1500 Ok(stored) => Ok(Some(stored)),
1501 Err(error) => Err(error),
1502 })
1503 })
1504 .with_context(|| format!("clearing the MCP OAuth token for '{server_name}'"))?
1505 }
1506
1507 fn parse_stored_oauth_tokens(serialized: &str, server_name: &str) -> Result<StoredMcpOAuthTokens> {
1508 serde_json::from_str(serialized).map_err(|_| {
1509 anyhow!(
1510 "stored MCP OAuth token for '{server_name}' is not valid credential JSON; contents were omitted"
1511 )
1512 })
1513 }
1514
1515 pub(crate) fn save_oauth_tokens(tokens: &StoredMcpOAuthTokens) -> Result<()> {
1516 let secrets = codewhale_secrets::Secrets::auto_detect();
1517 let key = store_key(&tokens.server_name, &tokens.url);
1518 let serialized = serde_json::to_string(tokens).context("serializing MCP OAuth token")?;
1519 secrets
1520 .set(&key, &serialized)
1521 .with_context(|| format!("saving MCP OAuth token for '{}'", tokens.server_name))
1522 }
1523
1524 fn delete_oauth_tokens(server_name: &str, url: &str) -> Result<bool> {
1525 let secrets = codewhale_secrets::Secrets::auto_detect();
1526 let key = store_key(server_name, url);
1527 let existed = secrets
1528 .get(&key)
1529 .with_context(|| format!("reading MCP OAuth token for '{server_name}'"))?
1530 .is_some();
1531 secrets
1532 .delete(&key)
1533 .with_context(|| format!("deleting MCP OAuth token for '{server_name}'"))?;
1534 Ok(existed)
1535 }
1536
1537 fn store_key(server_name: &str, url: &str) -> String {
1538 let mut payload = Vec::with_capacity(server_name.len() + url.len() + 1);
1539 payload.extend_from_slice(server_name.as_bytes());
1540 payload.push(0);
1541 payload.extend_from_slice(url.as_bytes());
1542 let digest = Sha256::digest(&payload);
1543 format!("mcp_oauth_{}", URL_SAFE_NO_PAD.encode(digest))
1544 }
1545
1546 fn oauth_tokens_are_usable(tokens: &StoredMcpOAuthTokens) -> bool {
1547 if tokens.client_id.trim().is_empty() {
1548 return false;
1549 }
1550 let response = &tokens.token_response.0;
1551 if token_needs_refresh(tokens.expires_at) {
1552 return response
1553 .refresh_token()
1554 .is_some_and(|token| !token.secret().trim().is_empty());
1555 }
1556 !response.access_token().secret().trim().is_empty()
1557 }
1558
1559 fn refresh_expires_in_from_timestamp(tokens: &mut StoredMcpOAuthTokens) {
1560 let Some(expires_at) = tokens.expires_at else {
1561 return;
1562 };
1563 match expires_in_from_timestamp(expires_at) {
1564 Some(seconds) => {
1565 let duration = Duration::from_secs(seconds);
1566 tokens.token_response.0.set_expires_in(Some(&duration));
1567 }
1568 None => {
1569 tokens
1570 .token_response
1571 .0
1572 .set_expires_in(Some(&Duration::ZERO));
1573 }
1574 }
1575 }
1576
1577 fn compute_expires_at_millis(response: &OAuthTokenResponse) -> Option<u64> {
1578 let expires = response.expires_in()?;
1579 let now = SystemTime::now()
1580 .duration_since(UNIX_EPOCH)
1581 .ok()?
1582 .as_millis() as u64;
1583 Some(now.saturating_add(expires.as_millis() as u64))
1584 }
1585
1586 fn expires_in_from_timestamp(expires_at: u64) -> Option<u64> {
1587 let now = SystemTime::now()
1588 .duration_since(UNIX_EPOCH)
1589 .ok()?
1590 .as_millis() as u64;
1591 if expires_at <= now {
1592 return None;
1593 }
1594 Some((expires_at - now) / 1000)
1595 }
1596
1597 fn token_needs_refresh(expires_at: Option<u64>) -> bool {
1598 let Some(expires_at) = expires_at else {
1599 return false;
1600 };
1601 let now = SystemTime::now()
1602 .duration_since(UNIX_EPOCH)
1603 .map(|duration| duration.as_millis() as u64)
1604 .unwrap_or(0);
1605 now.saturating_add(REFRESH_SKEW_MILLIS) >= expires_at
1606 }
1607
1608 struct CallbackServerGuard {
1609 accept_task: tokio::task::JoinHandle<()>,
1610 }
1611
1612 impl Drop for CallbackServerGuard {
1613 fn drop(&mut self) {
1614 // Aborting drops the accept future and its owned listener instead of
1615 // leaving a detached task holding a fixed callback port indefinitely.
1616 self.accept_task.abort();
1617 }
1618 }
1619
1620 struct OauthLoginFlow {
1621 auth_url: String,
1622 oauth_state: OAuthState,
1623 rx: oneshot::Receiver<CallbackResult>,
1624 guard: CallbackServerGuard,
1625 server_name: String,
1626 server_url: String,
1627 }
1628
1629 impl OauthLoginFlow {
1630 #[allow(clippy::too_many_arguments)]
1631 async fn new(
1632 server_name: &str,
1633 server_url: &str,
1634 client: McpHttpClient,
1635 scopes: &[String],
1636 oauth_client_id: Option<&str>,
1637 oauth_resource: Option<&str>,
1638 callback_port: Option<u16>,
1639 callback_url: Option<&str>,
1640 ) -> Result<Self> {
1641 let bind_host = callback_bind_host(callback_url);
1642 let bind_addr = match callback_port {
1643 Some(0) => bail!("invalid MCP OAuth callback port 0"),
1644 Some(port) => format!("{bind_host}:{port}"),
1645 None => format!("{bind_host}:0"),
1646 };
1647 let listener = TcpListener::bind(&bind_addr)
1648 .await
1649 .map_err(|err| anyhow!(err))?;
1650 let redirect_uri = resolve_redirect_uri(&listener, callback_url)?;
1651 let callback_id = callback_id_from_server_url(server_url)?;
1652 let redirect_uri = append_callback_id_to_redirect_uri(&redirect_uri, &callback_id)?;
1653 let callback_path = callback_path_from_redirect_uri(&redirect_uri)?;
1654
1655 let (tx, rx) = oneshot::channel();
1656 let guard = CallbackServerGuard {
1657 accept_task: spawn_callback_server(listener, tx, callback_path),
1658 };
1659
1660 let scope_refs: Vec<&str> = scopes.iter().map(String::as_str).collect();
1661 let oauth_state = start_authorization(
1662 server_url,
1663 client,
1664 &scope_refs,
1665 &redirect_uri,
1666 oauth_client_id,
1667 )
1668 .await?;
1669 let auth_url = append_query_param(
1670 &oauth_state.get_authorization_url().await?,
1671 "resource",
1672 oauth_resource,
1673 );
1674 // #6040: logout clears this machine's token only — the provider keeps
1675 // its standing grant. Without a forced prompt the next login silently
1676 // re-grants it (same account/workspace, no picker ever shown), so an
1677 // explicit login could never change the authorized workspace.
1678 let auth_url = append_query_param(&auth_url, "prompt", Some("consent"));
1679
1680 Ok(Self {
1681 auth_url,
1682 oauth_state,
1683 rx,
1684 guard,
1685 server_name: server_name.to_string(),
1686 server_url: server_url.to_string(),
1687 })
1688 }
1689
1690 async fn finish(self) -> Result<()> {
1691 self.finish_with_announce(OAuthLoginAnnounce::Terminal)
1692 .await
1693 }
1694
1695 async fn finish_with_announce(mut self, announce: OAuthLoginAnnounce) -> Result<()> {
1696 match announce {
1697 OAuthLoginAnnounce::Terminal => {
1698 println!(
1699 "Authorize `{}` by opening this URL in your browser:\n{}\n",
1700 self.server_name, self.auth_url
1701 );
1702 if webbrowser::open(&self.auth_url).is_err() {
1703 eprintln!("Browser launch failed; copy the URL above manually.");
1704 }
1705 println!(
1706 "Waiting for browser authorization for MCP server '{}'...",
1707 self.server_name
1708 );
1709 }
1710 OAuthLoginAnnounce::Tool { open_browser } => {
1711 // A tool call is not a terminal: nothing goes to stdout. The
1712 // tool result carries the URL for the model to relay; the
1713 // browser open is a best-effort convenience on top.
1714 if open_browser {
1715 let _ = webbrowser::open(&self.auth_url);
1716 }
1717 }
1718 }
1719
1720 let result = async {
1721 let callback = timeout(Duration::from_secs(300), &mut self.rx)
1722 .await
1723 .with_context(|| {
1724 let retry_hint = match announce {
1725 OAuthLoginAnnounce::Terminal => "Retry from a terminal, or use task_shell_start/background shell if an agent is running the login flow.".to_string(),
1726 OAuthLoginAnnounce::Tool { .. } => format!(
1727 "The user can complete the sign-in directly via `/mcp login {}` or `codewhale mcp login {}`, then this tool can be called again.",
1728 self.server_name, self.server_name
1729 ),
1730 };
1731 format!(
1732 "timed out waiting for OAuth callback for MCP server '{}'. {retry_hint}",
1733 self.server_name
1734 )
1735 })?
1736 .context("OAuth callback was cancelled")?;
1737 let OauthCallbackResult {
1738 code,
1739 state,
1740 issuer,
1741 } = match callback {
1742 CallbackResult::Success(callback) => callback,
1743 CallbackResult::Error(error) => return Err(anyhow!(error)),
1744 };
1745
1746 // RFC 9207: servers that advertise
1747 // `authorization_response_iss_parameter_supported` send `iss` on the
1748 // redirect and rmcp requires it back; forward it so the callback binds
1749 // to the discovered issuer instead of failing as "missing".
1750 self.oauth_state
1751 .handle_callback_with_issuer(&code, &state, issuer.as_deref())
1752 .await
1753 .context("handling MCP OAuth callback")?;
1754
1755 let (client_id, credentials) = self
1756 .oauth_state
1757 .get_credentials()
1758 .await
1759 .context("reading MCP OAuth credentials")?;
1760 let credentials =
1761 credentials.ok_or_else(|| anyhow!("OAuth provider did not return credentials"))?;
1762 let stored = StoredMcpOAuthTokens {
1763 server_name: self.server_name.clone(),
1764 url: self.server_url.clone(),
1765 client_id,
1766 expires_at: compute_expires_at_millis(&credentials),
1767 token_response: WrappedOAuthTokenResponse(credentials),
1768 };
1769 save_oauth_tokens(&stored)
1770 }
1771 .await;
1772
1773 drop(self.guard);
1774 result
1775 }
1776 }
1777
1778 async fn start_authorization(
1779 server_url: &str,
1780 client: McpHttpClient,
1781 scopes: &[&str],
1782 redirect_uri: &str,
1783 oauth_client_id: Option<&str>,
1784 ) -> Result<OAuthState> {
1785 let Some(client_id) = oauth_client_id.filter(|client_id| !client_id.trim().is_empty()) else {
1786 let mut attempt_scopes: Vec<String> =
1787 scopes.iter().map(|scope| (*scope).to_string()).collect();
1788 // Dynamic registration may reject part of the scope list the server
1789 // itself advertised (Supabase validates registration scopes against a
1790 // narrower allow-list than its `scopes_supported`). Drop exactly the
1791 // scopes the server named invalid and retry once; if it named none,
1792 // register without scopes so the server applies its defaults.
1793 for retried in [false, true] {
1794 let mut oauth_state = OAuthState::new_with_oauth_http_client(
1795 server_url,
1796 Arc::new(RecordingOAuthHttpClient::new(client.clone())),
1797 )
1798 .await?;
1799 let started = oauth_state
1800 .start_authorization(
1801 AuthorizationRequest::new(redirect_uri)
1802 .with_scopes(attempt_scopes.iter().map(String::as_str))
1803 .with_client_name("Codewhale"),
1804 )
1805 .await;
1806 match started {
1807 Ok(()) => return Ok(oauth_state),
1808 Err(error) if !retried && !attempt_scopes.is_empty() => {
1809 let message = error.to_string();
1810 let Some(narrowed) =
1811 scopes_after_registration_rejection(&attempt_scopes, &message)
1812 else {
1813 return Err(error.into());
1814 };
1815 tracing::warn!(
1816 target: "mcp::oauth",
1817 server_url,
1818 dropped = attempt_scopes.len() - narrowed.len(),
1819 "OAuth client registration rejected part of the requested scope list; retrying with the accepted scopes"
1820 );
1821 attempt_scopes = narrowed;
1822 }
1823 Err(error) => return Err(error.into()),
1824 }
1825 }
1826 unreachable!("registration retry loop returns on success or error");
1827 };
1828
1829 let mut manager = AuthorizationManager::new_with_oauth_http_client(
1830 server_url,
1831 Arc::new(RecordingOAuthHttpClient::new(client)),
1832 )
1833 .await?;
1834 let metadata = manager.resolve_metadata().await?.metadata;
1835 manager.set_metadata(metadata);
1836 manager.configure_client(
1837 OAuthClientConfig::new(client_id, redirect_uri)
1838 .with_scopes(scopes.iter().map(|scope| (*scope).to_string()).collect()),
1839 )?;
1840 let auth_url = manager.get_authorization_url(scopes).await?;
1841 Ok(OAuthState::Session(
1842 AuthorizationSession::for_scope_upgrade(manager, auth_url, redirect_uri),
1843 ))
1844 }
1845
1846 /// Given a registration failure message, return the scopes to retry with, or
1847 /// `None` when the failure is not about scopes. Servers that validate the
1848 /// `scope` field report positions like `scope.3: Invalid option`; those exact
1849 /// entries are dropped. A scope error without positions retries with no
1850 /// scopes at all, letting the server grant its defaults.
1851 fn scopes_after_registration_rejection(scopes: &[String], message: &str) -> Option<Vec<String>> {
1852 let lower = message.to_ascii_lowercase();
1853 if !(lower.contains("registration") && lower.contains("scope")) {
1854 return None;
1855 }
1856 let mut rejected = std::collections::BTreeSet::new();
1857 for (start, _) in message.match_indices("scope.") {
1858 let digits: String = message[start + "scope.".len()..]
1859 .chars()
1860 .take_while(char::is_ascii_digit)
1861 .collect();
1862 if let Ok(index) = digits.parse::<usize>() {
1863 rejected.insert(index);
1864 }
1865 }
1866 let narrowed: Vec<String> = scopes
1867 .iter()
1868 .enumerate()
1869 .filter(|(index, _)| !rejected.contains(index))
1870 .map(|(_, scope)| scope.clone())
1871 .collect();
1872 if narrowed.len() == scopes.len() {
1873 // The server complained about scopes without naming any position:
1874 // the only safe retry is to omit the field.
1875 return Some(Vec::new());
1876 }
1877 Some(narrowed)
1878 }
1879
1880 fn spawn_callback_server(
1881 listener: TcpListener,
1882 tx: oneshot::Sender<CallbackResult>,
1883 expected_callback_path: String,
1884 ) -> tokio::task::JoinHandle<()> {
1885 tokio::spawn(async move {
1886 // The sender is wrapped in Option so we can take it on success/error
1887 let mut tx_opt = Some(tx);
1888 loop {
1889 let (mut stream, _) = match listener.accept().await {
1890 Ok(pair) => pair,
1891 Err(_) => break,
1892 };
1893 let path = match read_http_path(&mut stream).await {
1894 Some(p) => p,
1895 None => {
1896 let _ = write_http_response(&mut stream, 400, "Invalid OAuth callback").await;
1897 continue;
1898 }
1899 };
1900 match parse_oauth_callback(&path, &expected_callback_path) {
1901 CallbackOutcome::Success(callback) => {
1902 let _ = write_http_response(
1903 &mut stream,
1904 200,
1905 "Authentication complete. You may close this window.",
1906 )
1907 .await;
1908 if let Some(tx) = tx_opt.take() {
1909 let _ = tx.send(CallbackResult::Success(callback));
1910 }
1911 break;
1912 }
1913 CallbackOutcome::Error(error) => {
1914 let msg = error.to_string();
1915 let _ = write_http_response(&mut stream, 400, &msg).await;
1916 if let Some(tx) = tx_opt.take() {
1917 let _ = tx.send(CallbackResult::Error(error));
1918 }
1919 break;
1920 }
1921 CallbackOutcome::Invalid => {
1922 let _ = write_http_response(&mut stream, 400, "Invalid OAuth callback").await;
1923 }
1924 }
1925 }
1926 })
1927 }
1928
1929 async fn read_http_path(stream: &mut tokio::net::TcpStream) -> Option<String> {
1930 let mut buf = Vec::new();
1931 let mut tmp = [0u8; 1024];
1932 // Read until we have \r\n\r\n or exceed limit
1933 loop {
1934 match stream.read(&mut tmp).await {
1935 Ok(0) => break,
1936 Ok(n) => {
1937 buf.extend_from_slice(&tmp[..n]);
1938 if buf.windows(4).any(|w| w == b"\r\n\r\n") {
1939 break;
1940 }
1941 if buf.len() > 8192 {
1942 break;
1943 }
1944 }
1945 Err(_) => return None,
1946 }
1947 }
1948 let request = String::from_utf8_lossy(&buf);
1949 let first_line = request.lines().next()?;
1950 // Expected: GET /callback?code=... HTTP/1.1
1951 let mut parts = first_line.split_whitespace();
1952 let _method = parts.next()?;
1953 let path = parts.next()?.to_string();
1954 Some(path)
1955 }
1956
1957 async fn write_http_response(
1958 stream: &mut tokio::net::TcpStream,
1959 status: u16,
1960 body: &str,
1961 ) -> std::io::Result<()> {
1962 let status_text = match status {
1963 200 => "OK",
1964 400 => "Bad Request",
1965 _ => "OK",
1966 };
1967 let response = format!(
1968 "HTTP/1.1 {status} {status_text}\r\nContent-Length: {}\r\nContent-Type: text/plain\r\nConnection: close\r\n\r\n{body}",
1969 body.len()
1970 );
1971 stream.write_all(response.as_bytes()).await?;
1972 stream.flush().await?;
1973 Ok(())
1974 }
1975
1976 #[derive(Debug, Clone, PartialEq, Eq)]
1977 struct OauthCallbackResult {
1978 code: String,
1979 state: String,
1980 /// RFC 9207 `iss` from the redirect, when the authorization server sends it.
1981 issuer: Option<String>,
1982 }
1983
1984 enum CallbackResult {
1985 Success(OauthCallbackResult),
1986 Error(OAuthProviderError),
1987 }
1988
1989 #[derive(Debug, Clone, PartialEq, Eq)]
1990 enum CallbackOutcome {
1991 Success(OauthCallbackResult),
1992 Error(OAuthProviderError),
1993 Invalid,
1994 }
1995
1996 fn parse_oauth_callback(path: &str, expected_callback_path: &str) -> CallbackOutcome {
1997 let Some((route, query)) = path.split_once('?') else {
1998 return CallbackOutcome::Invalid;
1999 };
2000 if route != expected_callback_path {
2001 return CallbackOutcome::Invalid;
2002 }
2003
2004 let mut code = None;
2005 let mut state = None;
2006 let mut issuer = None;
2007 let mut error = None;
2008 let mut error_description = None;
2009 for pair in query.split('&') {
2010 let Some((key, value)) = pair.split_once('=') else {
2011 continue;
2012 };
2013 let Ok(decoded) = decode(value) else {
2014 continue;
2015 };
2016 let decoded = decoded.into_owned();
2017 match key {
2018 "code" => code = Some(decoded),
2019 "state" => state = Some(decoded),
2020 "iss" => issuer = Some(decoded),
2021 "error" => error = Some(decoded),
2022 "error_description" => error_description = Some(decoded),
2023 _ => {}
2024 }
2025 }
2026
2027 if let (Some(code), Some(state)) = (code, state) {
2028 return CallbackOutcome::Success(OauthCallbackResult {
2029 code,
2030 state,
2031 issuer,
2032 });
2033 }
2034 if error.is_some() || error_description.is_some() {
2035 return CallbackOutcome::Error(OAuthProviderError::new(error, error_description));
2036 }
2037 CallbackOutcome::Invalid
2038 }
2039
2040 fn local_redirect_uri(listener: &TcpListener) -> Result<String> {
2041 let addr = listener.local_addr()?;
2042 match addr {
2043 std::net::SocketAddr::V4(v4) => Ok(format!("http://{}:{}/callback", v4.ip(), v4.port())),
2044 std::net::SocketAddr::V6(v6) => Ok(format!("http://[{}]:{}/callback", v6.ip(), v6.port())),
2045 }
2046 }
2047
2048 fn resolve_redirect_uri(listener: &TcpListener, callback_url: Option<&str>) -> Result<String> {
2049 let Some(callback_url) = callback_url else {
2050 return local_redirect_uri(listener);
2051 };
2052 Url::parse(callback_url)
2053 .with_context(|| format!("invalid MCP OAuth callback URL '{callback_url}'"))?;
2054 Ok(callback_url.to_string())
2055 }
2056
2057 fn callback_bind_host(callback_url: Option<&str>) -> &'static str {
2058 let Some(callback_url) = callback_url else {
2059 return "127.0.0.1";
2060 };
2061 let Ok(parsed) = Url::parse(callback_url) else {
2062 return "127.0.0.1";
2063 };
2064 match parsed.host_str() {
2065 Some("localhost" | "127.0.0.1" | "::1") | None => "127.0.0.1",
2066 Some(_) => "0.0.0.0",
2067 }
2068 }
2069
2070 fn callback_id_from_server_url(server_url: &str) -> Result<String> {
2071 let mut parsed =
2072 Url::parse(server_url).with_context(|| format!("invalid MCP server URL '{server_url}'"))?;
2073 parsed
2074 .host_str()
2075 .ok_or_else(|| anyhow!("MCP server URL '{server_url}' must include a host"))?;
2076 parsed.set_fragment(None);
2077 let digest = Sha256::digest(parsed.as_str().as_bytes());
2078 Ok(URL_SAFE_NO_PAD.encode(&digest[..9]))
2079 }
2080
2081 fn append_callback_id_to_redirect_uri(redirect_uri: &str, callback_id: &str) -> Result<String> {
2082 let mut parsed = Url::parse(redirect_uri)
2083 .with_context(|| format!("invalid redirect URI '{redirect_uri}'"))?;
2084 let path = parsed.path();
2085 let new_path = if path.ends_with('/') {
2086 format!("{path}{callback_id}")
2087 } else {
2088 format!("{path}/{callback_id}")
2089 };
2090 parsed.set_path(&new_path);
2091 Ok(parsed.to_string())
2092 }
2093
2094 fn callback_path_from_redirect_uri(redirect_uri: &str) -> Result<String> {
2095 let parsed = Url::parse(redirect_uri)
2096 .with_context(|| format!("invalid redirect URI '{redirect_uri}'"))?;
2097 Ok(parsed.path().to_string())
2098 }
2099
2100 fn append_query_param(url: &str, key: &str, value: Option<&str>) -> String {
2101 let Some(value) = value else {
2102 return url.to_string();
2103 };
2104 let value = value.trim();
2105 if value.is_empty() {
2106 return url.to_string();
2107 }
2108 if let Ok(mut parsed) = Url::parse(url) {
2109 parsed.query_pairs_mut().append_pair(key, value);
2110 return parsed.to_string();
2111 }
2112 let separator = if url.contains('?') { "&" } else { "?" };
2113 format!("{url}{separator}{key}={}", urlencoding::encode(value))
2114 }
2115
2116 impl McpServerConfig {
2117 pub fn oauth_client_id(&self) -> Option<&str> {
2118 self.oauth
2119 .as_ref()
2120 .and_then(|oauth| oauth.client_id.as_deref())
2121 }
2122 }
2123
2124 #[cfg(test)]
2125 mod tests {
2126 #[test]
2127 fn typed_oauth_client_requires_https_or_explicit_loopback_without_header_guessing() {
2128 crate::tls::ensure_rustls_crypto_provider();
2129 let server: super::McpServerConfig = serde_json::from_value(serde_json::json!({
2130 "url": "https://example.invalid/mcp"
2131 }))
2132 .unwrap();
2133 let error = super::oauth_http_client(&server, "http://example.invalid/mcp", None)
2134 .err()
2135 .expect("OAuth form credentials must not use public HTTP");
2136 assert!(error.to_string().contains("credentials require HTTPS"));
2137 assert!(super::oauth_http_client(&server, "https://example.invalid/mcp", None).is_ok());
2138 assert!(super::oauth_http_client(&server, "http://127.0.0.1/mcp", None).is_ok());
2139 assert!(
2140 super::oauth_http_client(
2141 &server,
2142 "https://fixture-user:fixture-password@example.invalid/mcp",
2143 None
2144 )
2145 .is_err()
2146 );
2147 }
2148
2149 #[test]
2150 fn registration_rejection_drops_exactly_the_named_scopes() {
2151 let scopes: Vec<String> = ["a", "b", "c", "d"].iter().map(|s| s.to_string()).collect();
2152 let message = concat!(
2153 "Registration failed: Dynamic registration failed: HTTP 400 Bad Request: ",
2154 "{\"message\":\"scope.1: Invalid option: expected one of \\\"a\\\"|\\\"c\\\",",
2155 "scope.3: Invalid option\"}"
2156 );
2157 assert_eq!(
2158 super::scopes_after_registration_rejection(&scopes, message),
2159 Some(vec!["a".to_string(), "c".to_string()])
2160 );
2161 // A scope complaint without positions retries without scopes.
2162 assert_eq!(
2163 super::scopes_after_registration_rejection(
2164 &scopes,
2165 "Registration failed: invalid scope"
2166 ),
2167 Some(Vec::new())
2168 );
2169 // Unrelated registration failures are not retried.
2170 assert_eq!(
2171 super::scopes_after_registration_rejection(&scopes, "Registration failed: HTTP 500"),
2172 None
2173 );
2174 assert_eq!(
2175 super::scopes_after_registration_rejection(&scopes, "network unreachable"),
2176 None
2177 );
2178 }
2179
2180 #[test]
2181 fn a_refresh_parse_failure_names_the_login_remedy_and_the_server() {
2182 let text = super::refresh_failure_context("supabase", true, None);
2183 assert!(text.contains("server supabase"));
2184 assert!(text.contains("codewhale mcp login supabase"), "{text}");
2185 assert!(text.contains("/mcp login supabase"), "{text}");
2186 assert!(!text.contains("it answered"), "{text}");
2187 // An auth-required failure keeps the plain context: the typed state
2188 // and the login tool already carry the remedy.
2189 let plain = super::refresh_failure_context("supabase", false, None);
2190 assert_eq!(plain, "refreshing MCP OAuth token for server supabase");
2191 }
2192
2193 #[test]
2194 fn a_refresh_parse_failure_keeps_the_token_endpoints_receipt() {
2195 // The supabase receipt (#5926): rmcp said only "Failed to parse
2196 // server response". With the status line and a masked excerpt the
2197 // operator can tell a provider's HTML 502 from our parser.
2198 let mut headers = HeaderMap::new();
2199 headers.insert(
2200 CONTENT_TYPE,
2201 HeaderValue::from_static("text/html; charset=utf-8"),
2202 );
2203 let receipt = TokenEndpointReceipt::from_response(
2204 reqwest::StatusCode::BAD_GATEWAY,
2205 &headers,
2206 b"<html>\n <body>502 Bad Gateway</body>\n</html>",
2207 );
2208 let text = super::refresh_failure_context("supabase", true, Some(&receipt));
2209 assert!(
2210 text.contains(
2211 "it answered HTTP 502 Bad Gateway (text/html; charset=utf-8): <html> <body>502 Bad Gateway</body> </html>"
2212 ),
2213 "{text}"
2214 );
2215 assert!(text.contains("codewhale mcp login supabase"), "{text}");
2216 }
2217
2218 #[test]
2219 fn a_token_receipt_masks_credentials_before_it_cuts_the_body() {
2220 let secret = "sk-live-0123456789abcdef";
2221 let long_tail = "x".repeat(400);
2222 let body = format!(
2223 "{{\"token_type\":\"Bearer\",\"access_token\":\"{secret}\",\"refresh_token\": \"{secret}-r\",\"id_token\":\"{secret}-id\",\"expires_in\":\"soon\",\"note\":\"{long_tail}\"}}"
2224 );
2225 let mut headers = HeaderMap::new();
2226 headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
2227 let receipt =
2228 TokenEndpointReceipt::from_response(reqwest::StatusCode::OK, &headers, body.as_bytes());
2229 let text = receipt.to_string();
2230 assert!(!text.contains(secret), "{text}");
2231 assert!(text.contains("\"access_token\":\"***\""), "{text}");
2232 assert!(text.contains("\"refresh_token\": \"***\""), "{text}");
2233 assert!(text.contains("\"id_token\":\"***\""), "{text}");
2234 // Non-secret members survive so the shape of the answer is readable.
2235 assert!(text.contains("\"expires_in\":\"soon\""), "{text}");
2236 assert!(text.contains("\"token_type\":\"Bearer\""), "{text}");
2237 assert!(text.ends_with('…'), "{text}");
2238 assert!(receipt.excerpt.len() <= TOKEN_RECEIPT_EXCERPT_BYTES + '…'.len_utf8());
2239 }
2240
2241 #[test]
2242 fn oauth_secret_masking_covers_form_pairs_bearer_schemes_and_case() {
2243 assert_eq!(
2244 mask_oauth_secrets("client_secret=abc123&grant_type=refresh_token&refresh_token=zzz"),
2245 "client_secret=***&grant_type=refresh_token&refresh_token=***"
2246 );
2247 assert_eq!(
2248 mask_oauth_secrets("Authorization: Bearer eyJhbGciOi.payload.sig, retry"),
2249 "Authorization: ***, retry"
2250 );
2251 assert_eq!(
2252 mask_oauth_secrets("{\"Access_Token\": \"quoted \\\" inside\", \"scope\": \"read\"}"),
2253 "{\"Access_Token\": \"***\", \"scope\": \"read\"}"
2254 );
2255 // A field name that merely contains a secret name is not a secret.
2256 assert_eq!(
2257 mask_oauth_secrets("{\"error_code\":\"invalid_request\",\"my_access_token_count\":3}"),
2258 "{\"error_code\":\"invalid_request\",\"my_access_token_count\":3}"
2259 );
2260 // Multi-byte text around a secret stays intact.
2261 assert_eq!(
2262 mask_oauth_secrets("トークン access_token=秘密 終わり"),
2263 "トークン access_token=*** 終わり"
2264 );
2265 }
2266
2267 #[test]
2268 fn a_token_receipt_names_an_empty_body_and_a_missing_content_type() {
2269 let receipt = TokenEndpointReceipt::from_response(
2270 reqwest::StatusCode::SERVICE_UNAVAILABLE,
2271 &HeaderMap::new(),
2272 b"",
2273 );
2274 assert_eq!(
2275 receipt.to_string(),
2276 "HTTP 503 Service Unavailable (no content-type) with an empty body"
2277 );
2278 }
2279
2280 use super::*;
2281 use std::sync::atomic::{AtomicBool, Ordering};
2282
2283 #[test]
2284 fn stored_credential_identity_ignores_only_a_derived_expiry_countdown() {
2285 let original = serde_json::json!({
2286 "server_name": "fixture", "url": "https://example.invalid/mcp",
2287 "client_id": "fixture-client", "expires_at": 9_999_999_999_999_u64,
2288 "token_response": {
2289 "access_token": "fixture-access", "refresh_token": "fixture-refresh",
2290 "token_type": "Bearer", "expires_in": 3600, "scope": "read"
2291 }
2292 });
2293 let held: StoredMcpOAuthTokens = serde_json::from_value(original.clone()).unwrap();
2294 let mut aged = original.clone();
2295 aged["token_response"]["expires_in"] = serde_json::json!(3598);
2296 let loaded: StoredMcpOAuthTokens = serde_json::from_value(aged.clone()).unwrap();
2297 assert!(
2298 held == loaded,
2299 "elapsed time alone is not credential rotation"
2300 );
2301 for (field, value) in [
2302 ("access_token", "new-access"),
2303 ("refresh_token", "new-refresh"),
2304 ("scope", "read write"),
2305 ("token_type", "Mac"),
2306 ] {
2307 let mut rotated = aged.clone();
2308 rotated["token_response"][field] = serde_json::json!(value);
2309 let rotated: StoredMcpOAuthTokens = serde_json::from_value(rotated).unwrap();
2310 assert!(
2311 held != rotated,
2312 "a changed {field} remains a distinct credential"
2313 );
2314 }
2315 for (field, value) in [
2316 ("server_name", serde_json::json!("other")),
2317 ("client_id", serde_json::json!("other-client")),
2318 ("url", serde_json::json!("https://other.invalid/mcp")),
2319 ("expires_at", serde_json::json!(9_999_999_999_998_u64)),
2320 ] {
2321 let mut rotated = aged.clone();
2322 rotated[field] = value;
2323 let rotated: StoredMcpOAuthTokens = serde_json::from_value(rotated).unwrap();
2324 assert!(
2325 held != rotated,
2326 "a changed {field} remains a distinct credential"
2327 );
2328 }
2329 let mut legacy = held.clone();
2330 legacy.expires_at = None;
2331 let mut legacy_aged = loaded;
2332 legacy_aged.expires_at = None;
2333 assert!(
2334 legacy != legacy_aged,
2335 "without a durable deadline the stored lifetime is meaningful"
2336 );
2337 }
2338
2339 #[test]
2340 fn resolve_oauth_scopes_prefers_explicit() {
2341 let resolved = resolve_oauth_scopes(
2342 Some(vec!["explicit".to_string()]),
2343 vec!["configured".to_string()],
2344 Some(vec!["discovered".to_string()]),
2345 );
2346 assert_eq!(resolved.source, McpOAuthScopesSource::Explicit);
2347 assert_eq!(resolved.scopes, vec!["explicit"]);
2348 }
2349
2350 #[test]
2351 fn parse_oauth_callback_accepts_success() {
2352 let parsed = parse_oauth_callback("/callback/id?code=abc&state=xyz", "/callback/id");
2353 assert_eq!(
2354 parsed,
2355 CallbackOutcome::Success(OauthCallbackResult {
2356 code: "abc".to_string(),
2357 state: "xyz".to_string(),
2358 issuer: None,
2359 })
2360 );
2361 }
2362
2363 #[test]
2364 fn parse_oauth_callback_keeps_rfc9207_issuer() {
2365 // Cloudflare's MCP authorization server advertises
2366 // authorization_response_iss_parameter_supported and sends `iss` back;
2367 // dropping it makes rmcp reject the callback as missing a required issuer.
2368 let parsed = parse_oauth_callback(
2369 "/callback/id?code=abc&state=xyz&iss=https%3A%2F%2Fmcp.cloudflare.com",
2370 "/callback/id",
2371 );
2372 assert_eq!(
2373 parsed,
2374 CallbackOutcome::Success(OauthCallbackResult {
2375 code: "abc".to_string(),
2376 state: "xyz".to_string(),
2377 issuer: Some("https://mcp.cloudflare.com".to_string()),
2378 })
2379 );
2380 }
2381
2382 #[test]
2383 fn parse_oauth_callback_accepts_provider_error() {
2384 let parsed = parse_oauth_callback(
2385 "/callback/id?error=invalid_scope&error_description=nope",
2386 "/callback/id",
2387 );
2388 assert!(matches!(parsed, CallbackOutcome::Error(_)));
2389 }
2390
2391 #[test]
2392 fn store_key_does_not_include_raw_url_or_name() {
2393 let key = store_key("github", "https://example.com/mcp");
2394 assert!(key.starts_with("mcp_oauth_"));
2395 assert!(!key.contains("github"));
2396 assert!(!key.contains("example.com"));
2397 }
2398
2399 #[test]
2400 fn malformed_stored_oauth_diagnostic_omits_secret_contents_and_keys() {
2401 let secret = "cw-secret-mcp-oauth-4507";
2402 let serialized =
2403 format!(r#"{{"token_response":{{"access_token":"{secret}"}} trailing-junk}}"#);
2404 let error = parse_stored_oauth_tokens(&serialized, "private")
2405 .expect_err("malformed credential JSON must fail");
2406 let diagnostic = format!("{error:#}");
2407 assert!(!diagnostic.contains(secret), "{diagnostic}");
2408 assert!(!diagnostic.contains("access_token"), "{diagnostic}");
2409 assert!(diagnostic.contains("contents were omitted"), "{diagnostic}");
2410 }
2411
2412 #[test]
2413 fn auth_required_classifier_matches_http_401_shapes() {
2414 let err = anyhow!("MCP Streamable HTTP rejected status=401 Unauthorized");
2415 assert!(error_looks_auth_required(&err));
2416
2417 let err = anyhow!("authentication_required for remote server");
2418 assert!(error_looks_auth_required(&err));
2419
2420 let err = anyhow!("connection refused");
2421 assert!(!error_looks_auth_required(&err));
2422
2423 assert!(error_text_looks_auth_required("HTTP 401 from upstream"));
2424 assert!(error_text_looks_auth_required("request failed (401)"));
2425 // The failure the classifier used to misread: a reset on a loopback
2426 // server whose ephemeral port contains 401 is a transport error.
2427 let reset = "error sending request for url (http://127.0.0.1:50401/mcp): \
2428 client error (SendRequest): connection closed before message completed";
2429 assert!(!error_text_looks_auth_required(reset), "{reset}");
2430 assert!(!error_text_looks_auth_required(
2431 "read 14010 bytes before the stream reset"
2432 ));
2433 }
2434
2435 #[test]
2436 fn auth_required_classifier_treats_rejected_grants_as_auth_required() {
2437 // A definitively rejected refresh grant is recoverable only by a
2438 // fresh login, so it must classify like a 401 on every surface.
2439 let err = anyhow!("refreshing MCP OAuth token for server wiki")
2440 .context("Server returned error response: invalid_grant: stale grant");
2441 assert!(error_looks_auth_required(&err));
2442 assert!(error_text_looks_auth_required(
2443 "wiki requires OAuth — run /mcp login wiki"
2444 ));
2445 assert!(error_text_looks_auth_required("wiki: ◆ auth required"));
2446 assert!(!error_text_looks_auth_required(
2447 "invalid_request: missing parameter"
2448 ));
2449 // rmcp's own `AuthError::AuthorizationRequired` wording: a stored
2450 // credential that can no longer be refreshed is a login, not a
2451 // transport failure.
2452 assert!(error_text_looks_auth_required(
2453 "refreshing MCP OAuth token for server wiki: OAuth authorization required"
2454 ));
2455 assert!(!error_text_looks_auth_required(
2456 "authorization required for the requested file"
2457 ));
2458 }
2459
2460 #[test]
2461 fn auth_required_login_hint_names_server() {
2462 let hint = auth_required_login_hint("nordic-mcp");
2463 assert!(hint.contains("nordic-mcp"));
2464 assert!(hint.contains("codewhale mcp login nordic-mcp"));
2465 assert!(!hint.contains("/mcp auth"));
2466 }
2467
2468 #[test]
2469 fn tui_reauth_hints_name_the_login_command() {
2470 for hint in [tui_reauth_hint(), tui_reauth_refresh_failed_hint()] {
2471 assert!(
2472 hint.contains("/mcp login <name>"),
2473 "OAuth recovery must name the implemented command"
2474 );
2475 assert!(
2476 !hint.contains("/mcp auth"),
2477 "OAuth recovery must not advertise a missing /mcp auth command"
2478 );
2479 }
2480 assert!(error_text_looks_auth_required(
2481 "MCP server rejected the request with 401 Unauthorized"
2482 ));
2483 assert!(!error_text_looks_auth_required("connection refused"));
2484 }
2485
2486 #[tokio::test]
2487 async fn cancellable_oauth_drops_in_flight_flow_before_returning() {
2488 struct DropFlag(Arc<AtomicBool>);
2489 impl Drop for DropFlag {
2490 fn drop(&mut self) {
2491 self.0.store(true, Ordering::SeqCst);
2492 }
2493 }
2494
2495 let cancellation_token = CancellationToken::new();
2496 let cancel_from_task = cancellation_token.clone();
2497 let dropped = Arc::new(AtomicBool::new(false));
2498 let flow_dropped = Arc::clone(&dropped);
2499 let pending_flow = async move {
2500 let _guard = DropFlag(flow_dropped);
2501 std::future::pending::<Result<()>>().await
2502 };
2503 tokio::spawn(async move {
2504 tokio::task::yield_now().await;
2505 cancel_from_task.cancel();
2506 });
2507
2508 let error = run_cancellable_oauth(&cancellation_token, pending_flow)
2509 .await
2510 .expect_err("cancellation should stop the pending OAuth flow");
2511
2512 assert!(error.to_string().contains("OAuth login was cancelled"));
2513 assert!(
2514 dropped.load(Ordering::SeqCst),
2515 "the callback-server guard must be dropped before cancellation returns"
2516 );
2517 }
2518
2519 #[tokio::test]
2520 async fn callback_guard_aborts_accept_task_and_releases_fixed_port() -> Result<()> {
2521 let listener = TcpListener::bind(("127.0.0.1", 0)).await?;
2522 let addr = listener.local_addr()?;
2523 let (tx, _rx) = oneshot::channel();
2524 let guard = CallbackServerGuard {
2525 accept_task: spawn_callback_server(listener, tx, "/callback/test".to_string()),
2526 };
2527
2528 drop(guard);
2529
2530 let rebound = timeout(Duration::from_secs(1), async {
2531 loop {
2532 match TcpListener::bind(addr).await {
2533 Ok(listener) => break Ok(listener),
2534 Err(err) if err.kind() == std::io::ErrorKind::AddrInUse => {
2535 tokio::task::yield_now().await;
2536 }
2537 Err(err) => break Err(err),
2538 }
2539 }
2540 })
2541 .await
2542 .context("callback listener did not release its fixed port")??;
2543 drop(rebound);
2544 Ok(())
2545 }
2546
2547 async fn guarded_oauth_fixture(
2548 token_target: Option<String>,
2549 redirect_token: bool,
2550 ) -> (
2551 String,
2552 Arc<std::sync::atomic::AtomicUsize>,
2553 tokio::task::JoinHandle<()>,
2554 ) {
2555 use std::sync::atomic::{AtomicUsize, Ordering};
2556 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2557 let addr = listener.local_addr().unwrap();
2558 let captured = Arc::new(AtomicUsize::new(0));
2559 let seen = Arc::clone(&captured);
2560 let task = tokio::spawn(async move {
2561 loop {
2562 let Ok((mut socket, _)) = listener.accept().await else {
2563 break;
2564 };
2565 let mut bytes = Vec::new();
2566 let mut buffer = [0u8; 2048];
2567 loop {
2568 let n = socket.read(&mut buffer).await.unwrap();
2569 if n == 0 {
2570 break;
2571 }
2572 bytes.extend_from_slice(&buffer[..n]);
2573 if bytes.windows(4).any(|part| part == b"\r\n\r\n") {
2574 break;
2575 }
2576 }
2577 let request = String::from_utf8_lossy(&bytes);
2578 let path = request.split_whitespace().nth(1).unwrap_or("");
2579 let (status, extra, body) = if path == "/.well-known/oauth-authorization-server" {
2580 ("200 OK", String::new(), serde_json::json!({
2581 "issuer": format!("http://{addr}"),
2582 "authorization_endpoint": format!("http://{addr}/authorize"),
2583 "token_endpoint": token_target.clone().unwrap_or_else(|| format!("http://{addr}/token")),
2584 "registration_endpoint": format!("http://{addr}/register"),
2585 "response_types_supported": ["code"]
2586 }).to_string())
2587 } else if path == "/token" && redirect_token {
2588 (
2589 "307 Redirect",
2590 "Location: /capture\r\n".to_string(),
2591 String::new(),
2592 )
2593 } else if path == "/token" || path == "/capture" {
2594 seen.fetch_add(1, Ordering::SeqCst);
2595 ("200 OK", String::new(), r#"{"access_token":"new-fixture","token_type":"Bearer","refresh_token":"fixture-refresh"}"#.to_string())
2596 } else if path == "/register" {
2597 (
2598 "200 OK",
2599 String::new(),
2600 r#"{"client_id":"fixture-client","redirect_uris":[]}"#.to_string(),
2601 )
2602 } else {
2603 ("404 Not Found", String::new(), String::new())
2604 };
2605 let response = format!(
2606 "HTTP/1.1 {status}\r\n{extra}Content-Type: application/json\r\nConnection: close\r\nContent-Length: {}\r\n\r\n{body}",
2607 body.len()
2608 );
2609 let _ = socket.write_all(response.as_bytes()).await;
2610 }
2611 });
2612 (format!("http://{addr}/mcp"), captured, task)
2613 }
2614
2615 async fn guarded_oauth_state(
2616 url: &str,
2617 network_policy: Option<&NetworkPolicyDecider>,
2618 ) -> OAuthState {
2619 let client = McpHttpClient::new(
2620 url,
2621 false,
2622 false,
2623 false,
2624 network_policy,
2625 Duration::from_secs(1),
2626 Duration::from_secs(3),
2627 )
2628 .unwrap();
2629 let mut state = OAuthState::new_with_oauth_http_client(
2630 url,
2631 Arc::new(RecordingOAuthHttpClient::new(client)),
2632 )
2633 .await
2634 .unwrap();
2635 let tokens: OAuthTokenResponse = serde_json::from_value(serde_json::json!({
2636 "access_token":"fixture-access", "token_type":"Bearer", "refresh_token":"fixture-refresh"
2637 })).unwrap();
2638 state
2639 .set_credentials("fixture-client", tokens)
2640 .await
2641 .unwrap();
2642 state
2643 }
2644
2645 #[tokio::test]
2646 async fn guarded_oauth_refresh_honors_stop_and_preserves_normal_local_refresh() {
2647 use std::sync::atomic::Ordering;
2648 let _env = crate::test_support::lock_test_env();
2649 let _proxy = crate::test_support::EnvVarGuard::set("NO_PROXY", "*");
2650 for redirect in [true, false] {
2651 let (url, captured, task) = guarded_oauth_fixture(None, redirect).await;
2652 let state = guarded_oauth_state(&url, None).await;
2653 let result = state.refresh_token().await;
2654 if redirect {
2655 assert!(result.is_err(), "redirected refresh must not be followed");
2656 assert_eq!(captured.load(Ordering::SeqCst), 0);
2657 } else {
2658 result.unwrap();
2659 assert_eq!(captured.load(Ordering::SeqCst), 1);
2660 }
2661 task.abort();
2662 }
2663 }
2664
2665 #[tokio::test]
2666 async fn guarded_oauth_discovered_private_token_endpoint_never_receives_credentials() {
2667 let _env = crate::test_support::lock_test_env();
2668 let _proxy = crate::test_support::EnvVarGuard::set("NO_PROXY", "*");
2669 let destination = TcpListener::bind("127.0.0.1:0").await.unwrap();
2670 let target = format!("http://{}/token", destination.local_addr().unwrap());
2671 let (url, _, task) = guarded_oauth_fixture(Some(target), false).await;
2672 let state = guarded_oauth_state(&url, None).await;
2673 let error = tokio::time::timeout(Duration::from_secs(1), state.refresh_token())
2674 .await
2675 .expect("the destination guard rejects before attempting a network request")
2676 .unwrap_err();
2677 // rmcp intentionally wraps HTTP client failures as `Request failed`.
2678 // The observable invariant is an immediate failed refresh and no socket
2679 // at the private destination, rather than an SDK-specific error string.
2680 assert!(
2681 matches!(error, AuthError::TokenRefreshFailed(_)),
2682 "{error:#}"
2683 );
2684 assert!(
2685 tokio::time::timeout(Duration::from_millis(30), destination.accept())
2686 .await
2687 .is_err()
2688 );
2689 task.abort();
2690 }
2691
2692 #[tokio::test]
2693 async fn guarded_oauth_network_deny_applies_to_standalone_and_synthetic_login() {
2694 use crate::mcp::{AuthenticateToolStart, McpConfig, McpPool};
2695 use crate::network_policy::{DecisionToml, NetworkPolicy};
2696 let _env = crate::test_support::lock_test_env();
2697 let dir = tempfile::tempdir().unwrap();
2698 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
2699 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
2700 let _proxy = crate::test_support::EnvVarGuard::set("NO_PROXY", "*");
2701 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2702 let url = format!("http://{}/mcp", listener.local_addr().unwrap());
2703 let server: McpServerConfig =
2704 serde_json::from_value(serde_json::json!({"url":url})).unwrap();
2705 let denied = NetworkPolicyDecider::new(
2706 NetworkPolicy {
2707 default: DecisionToml::Deny,
2708 ..NetworkPolicy::default()
2709 },
2710 None,
2711 );
2712 let mut config = McpConfig::default();
2713 config
2714 .servers
2715 .insert("network-guard".to_string(), server.clone());
2716 let pool = McpPool::new(config).with_network_policy(denied.clone());
2717 let error = match pool.begin_authenticate_tool("network-guard").await {
2718 Err(error) => error,
2719 Ok(_) => panic!("a configured denied origin must not start synthetic authentication"),
2720 };
2721 assert!(error.to_string().contains("network policy"), "{error:#}");
2722 assert!(oauth_login_support(&server, Some(&denied)).await.is_err());
2723 assert_eq!(
2724 auth_status_for_server("network-guard", &server, Some(&denied)).await,
2725 McpAuthStatus::Unsupported
2726 );
2727 assert!(
2728 perform_oauth_login_for_server(
2729 "network-guard",
2730 &server,
2731 Some(vec!["explicit".to_string()]),
2732 None,
2733 None,
2734 Some(&denied)
2735 )
2736 .await
2737 .is_err()
2738 );
2739 assert!(
2740 tokio::time::timeout(Duration::from_millis(30), listener.accept())
2741 .await
2742 .is_err()
2743 );
2744
2745 let (url, _, task) = guarded_oauth_fixture(None, false).await;
2746 let server: McpServerConfig =
2747 serde_json::from_value(serde_json::json!({"url":url})).unwrap();
2748 let allowed = NetworkPolicyDecider::new(
2749 NetworkPolicy {
2750 default: DecisionToml::Allow,
2751 ..NetworkPolicy::default()
2752 },
2753 None,
2754 );
2755 let mut config = McpConfig::default();
2756 config.servers.insert("network-control".to_string(), server);
2757 let pool = McpPool::new(config).with_network_policy(allowed.clone());
2758 let AuthenticateToolStart::Login(login) = pool
2759 .begin_authenticate_tool("network-control")
2760 .await
2761 .unwrap()
2762 else {
2763 panic!("the configured local control must start a fresh login");
2764 };
2765 assert!(login.authorization_url().contains("/authorize"));
2766 // A later retry carries the same shared session ceiling.
2767 allowed.deny_session("127.0.0.1", "mcp");
2768 assert_eq!(
2769 login
2770 .network_policy
2771 .as_ref()
2772 .unwrap()
2773 .evaluate("127.0.0.1", "mcp"),
2774 crate::network_policy::Decision::Deny
2775 );
2776 drop(login);
2777 task.abort();
2778 }
2779
2780 #[tokio::test]
2781 async fn interactive_login_forces_the_consent_screen() {
2782 use crate::mcp::{AuthenticateToolStart, McpConfig, McpPool};
2783 use crate::network_policy::{DecisionToml, NetworkPolicy};
2784 let _env = crate::test_support::lock_test_env();
2785 let dir = tempfile::tempdir().unwrap();
2786 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
2787 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
2788 let _proxy = crate::test_support::EnvVarGuard::set("NO_PROXY", "*");
2789 let (url, _, task) = guarded_oauth_fixture(None, false).await;
2790 let server: McpServerConfig =
2791 serde_json::from_value(serde_json::json!({"url":url})).unwrap();
2792 let allowed = NetworkPolicyDecider::new(
2793 NetworkPolicy {
2794 default: DecisionToml::Allow,
2795 ..NetworkPolicy::default()
2796 },
2797 None,
2798 );
2799 let mut config = McpConfig::default();
2800 config.servers.insert("consent-probe".to_string(), server);
2801 let pool = McpPool::new(config).with_network_policy(allowed);
2802 let AuthenticateToolStart::Login(login) =
2803 pool.begin_authenticate_tool("consent-probe").await.unwrap()
2804 else {
2805 panic!("a fresh server must start an interactive login");
2806 };
2807
2808 // #6040: logout only clears this machine's token; the provider keeps
2809 // its standing grant, so the login URL must force the consent screen
2810 // or the same account/workspace is silently re-granted.
2811 // The URL carries the PKCE challenge and state, so the assertion
2812 // message reports only the fact that is being checked, never the URL.
2813 let forces_consent = login.authorization_url().contains("prompt=consent");
2814 assert!(
2815 forces_consent,
2816 "an interactive login must force consent (prompt=consent is missing from the authorization URL)"
2817 );
2818
2819 drop(login);
2820 task.abort();
2821 }
2822
2823 #[tokio::test]
2824 async fn guarded_oauth_refresh_keeps_live_session_network_denials() {
2825 use crate::network_policy::{DecisionToml, NetworkPolicy};
2826 let _env = crate::test_support::lock_test_env();
2827 let _proxy = crate::test_support::EnvVarGuard::set("NO_PROXY", "*");
2828 let (url, captured, task) = guarded_oauth_fixture(None, false).await;
2829 let policy = NetworkPolicyDecider::new(
2830 NetworkPolicy {
2831 default: DecisionToml::Allow,
2832 ..NetworkPolicy::default()
2833 },
2834 None,
2835 );
2836 let state = guarded_oauth_state(&url, Some(&policy)).await;
2837 state.refresh_token().await.unwrap();
2838 assert_eq!(captured.load(Ordering::SeqCst), 1);
2839 policy.deny_session("127.0.0.1", "mcp");
2840 assert!(state.refresh_token().await.is_err());
2841 assert_eq!(
2842 captured.load(Ordering::SeqCst),
2843 1,
2844 "no refresh request after session denial"
2845 );
2846 task.abort();
2847 }
2848 }
2849
2849 lines RUST