1use anyhow::{Context, Result, anyhow, bail};
4use base64::Engine;
5use base64::engine::general_purpose::URL_SAFE_NO_PAD;
6use reqwest::{Client, Url};
7use ring::rand::{SecureRandom, SystemRandom};
8use serde::{Deserialize, Serialize};
9use std::collections::BTreeMap;
10
11use crate::credentials::{AuthCredentialsStoreMode, CredentialStorage};
12use crate::pkce::{PkceChallenge, generate_pkce_challenge};
13
14const DEFAULT_CALLBACK_PORT: u16 = 8768;
15const DEFAULT_FLOW_TIMEOUT_SECS: u64 = 300;
16const REFRESH_SKEW_SECS: u64 = 60;
17
18#[derive(Debug, Clone, Serialize, Deserialize)]
20#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
21#[serde(default)]
22pub struct McpOAuthConfig {
23 pub authorization_url: String,
25 pub token_url: String,
27 pub client_id: String,
29 #[serde(default)]
31 pub scopes: Vec<String>,
32 #[serde(default)]
34 audience: Option<String>,
35 pub callback_port: u16,
37 flow_timeout_secs: u64,
39 #[serde(default)]
41 pub credentials_store_mode: AuthCredentialsStoreMode,
42 #[serde(default)]
44 extra_auth_params: BTreeMap<String, String>,
45 #[serde(default)]
47 extra_token_params: BTreeMap<String, String>,
48}
49
50impl Default for McpOAuthConfig {
51 fn default() -> Self {
52 Self {
53 authorization_url: String::new(),
54 token_url: String::new(),
55 client_id: String::new(),
56 scopes: Vec::new(),
57 audience: None,
58 callback_port: DEFAULT_CALLBACK_PORT,
59 flow_timeout_secs: DEFAULT_FLOW_TIMEOUT_SECS,
60 credentials_store_mode: AuthCredentialsStoreMode::default(),
61 extra_auth_params: BTreeMap::new(),
62 extra_token_params: BTreeMap::new(),
63 }
64 }
65}
66
67impl McpOAuthConfig {
68 fn validate(&self, provider_name: &str) -> Result<()> {
69 if self.authorization_url.trim().is_empty() {
70 bail!("MCP provider '{provider_name}' is missing oauth.authorization_url");
71 }
72 if self.token_url.trim().is_empty() {
73 bail!("MCP provider '{provider_name}' is missing oauth.token_url");
74 }
75 if self.client_id.trim().is_empty() {
76 bail!("MCP provider '{provider_name}' is missing oauth.client_id");
77 }
78 Ok(())
79 }
80
81 fn callback_url(&self) -> String {
82 format!("http://localhost:{}/auth/callback", self.callback_port)
83 }
84}
85
86#[derive(Debug, Clone, Serialize, Deserialize)]
88pub struct McpOAuthToken {
89 access_token: String,
90 refresh_token: Option<String>,
91 token_type: Option<String>,
92 scope: Option<String>,
93 obtained_at: u64,
94 expires_at: Option<u64>,
95}
96
97impl McpOAuthToken {
98 fn is_refresh_due(&self) -> bool {
99 self.expires_at
100 .is_some_and(|expires_at| now_secs().saturating_add(REFRESH_SKEW_SECS) >= expires_at)
101 }
102}
103
104#[derive(Debug, Clone)]
106pub enum McpOAuthStatus {
107 Authenticated { age_seconds: u64, expires_in: Option<u64> },
108 NotAuthenticated,
109}
110
111#[derive(Debug, Clone)]
113pub struct McpOAuthPreparedLogin {
114 pub auth_url: String,
115 pub callback_port: u16,
116 pub timeout_secs: u64,
117 pkce: PkceChallenge,
118 state: String,
119}
120
121impl McpOAuthPreparedLogin {
122 #[must_use]
123 pub fn expected_state(&self) -> &str {
124 &self.state
125 }
126}
127
128#[derive(Debug, Clone, PartialEq, Eq)]
130pub struct McpOAuthLoginCompletion {
131 pub name: String,
132 success: bool,
133 error: Option<String>,
134}
135
136#[derive(Debug, Clone, Default)]
138pub struct McpOAuthService;
139
140impl McpOAuthService {
141 #[must_use]
142 pub fn new() -> Self {
143 Self
144 }
145
146 pub fn prepare_login(&self, provider_name: &str, config: &McpOAuthConfig) -> Result<McpOAuthPreparedLogin> {
147 config.validate(provider_name)?;
148 let pkce = generate_pkce_challenge()?;
149 let state = generate_state()?;
150 let auth_url = build_auth_url(config, &pkce, &state)?;
151 Ok(McpOAuthPreparedLogin {
152 auth_url,
153 callback_port: config.callback_port,
154 timeout_secs: config.flow_timeout_secs,
155 pkce,
156 state,
157 })
158 }
159
160 pub async fn complete_login(
161 &self,
162 provider_name: &str,
163 config: &McpOAuthConfig,
164 prepared: &McpOAuthPreparedLogin,
165 code: &str,
166 ) -> Result<McpOAuthLoginCompletion> {
167 config.validate(provider_name)?;
168 let token = exchange_code_for_token(config, code, &prepared.pkce).await?;
169 save_token(provider_name, &token, config.credentials_store_mode)?;
170 Ok(McpOAuthLoginCompletion {
171 name: provider_name.to_string(),
172 success: true,
173 error: None,
174 })
175 }
176
177 pub fn status(&self, provider_name: &str, storage_mode: AuthCredentialsStoreMode) -> Result<McpOAuthStatus> {
178 let Some(token) = load_token(provider_name, storage_mode)? else {
179 return Ok(McpOAuthStatus::NotAuthenticated);
180 };
181 let now = now_secs();
182 Ok(McpOAuthStatus::Authenticated {
183 age_seconds: now.saturating_sub(token.obtained_at),
184 expires_in: token.expires_at.map(|expires_at| expires_at.saturating_sub(now)),
185 })
186 }
187
188 pub fn load_token(
189 &self,
190 provider_name: &str,
191 storage_mode: AuthCredentialsStoreMode,
192 ) -> Result<Option<McpOAuthToken>> {
193 load_token(provider_name, storage_mode)
194 }
195
196 pub async fn resolve_access_token(&self, provider_name: &str, config: &McpOAuthConfig) -> Result<Option<String>> {
197 let Some(mut token) = load_token(provider_name, config.credentials_store_mode)? else {
198 return Ok(None);
199 };
200
201 if token.is_refresh_due() {
202 if token.refresh_token.is_some() {
203 token = refresh_token(config, &token).await?;
204 save_token(provider_name, &token, config.credentials_store_mode)?;
205 } else {
206 bail!(
207 "Stored MCP OAuth token for '{provider_name}' expired and cannot be refreshed. Run `vtcode mcp login {provider_name}` again."
208 );
209 }
210 }
211
212 Ok(Some(token.access_token))
213 }
214
215 pub fn logout(
216 &self,
217 provider_name: &str,
218 storage_mode: AuthCredentialsStoreMode,
219 ) -> Result<McpOAuthLoginCompletion> {
220 clear_token(provider_name, storage_mode)?;
221 Ok(McpOAuthLoginCompletion {
222 name: provider_name.to_string(),
223 success: true,
224 error: None,
225 })
226 }
227}
228
229#[expect(
230 unused_results,
231 reason = "URL query builder methods mutate the serializer and return a fluent reference."
232)]
233fn build_auth_url(config: &McpOAuthConfig, challenge: &PkceChallenge, state: &str) -> Result<String> {
234 let mut url = Url::parse(&config.authorization_url).context("invalid oauth.authorization_url")?;
235 {
236 let mut query = url.query_pairs_mut();
237 query.append_pair("response_type", "code");
238 query.append_pair("client_id", &config.client_id);
239 query.append_pair("redirect_uri", &config.callback_url());
240 query.append_pair("code_challenge", &challenge.code_challenge);
241 query.append_pair("code_challenge_method", &challenge.code_challenge_method);
242 query.append_pair("state", state);
243 if !config.scopes.is_empty() {
244 query.append_pair("scope", &config.scopes.join(" "));
245 }
246 if let Some(audience) = config.audience.as_deref()
247 && !audience.trim().is_empty()
248 {
249 query.append_pair("audience", audience);
250 }
251 for (key, value) in &config.extra_auth_params {
252 if !key.trim().is_empty() {
253 query.append_pair(key, value);
254 }
255 }
256 }
257 Ok(url.to_string())
258}
259
260async fn exchange_code_for_token(
261 config: &McpOAuthConfig,
262 code: &str,
263 challenge: &PkceChallenge,
264) -> Result<McpOAuthToken> {
265 let mut form = vec![
266 ("grant_type".to_string(), "authorization_code".to_string()),
267 ("client_id".to_string(), config.client_id.clone()),
268 ("code".to_string(), code.to_string()),
269 ("redirect_uri".to_string(), config.callback_url()),
270 ("code_verifier".to_string(), challenge.code_verifier.to_string()),
271 ];
272 if let Some(audience) = config.audience.as_deref()
273 && !audience.trim().is_empty()
274 {
275 form.push(("audience".to_string(), audience.to_string()));
276 }
277 form.extend(
278 config
279 .extra_token_params
280 .iter()
281 .map(|(key, value)| (key.clone(), value.clone())),
282 );
283 send_token_request(&config.token_url, &form).await
284}
285
286async fn refresh_token(config: &McpOAuthConfig, current: &McpOAuthToken) -> Result<McpOAuthToken> {
287 let refresh_token = current
288 .refresh_token
289 .as_deref()
290 .filter(|value| !value.trim().is_empty())
291 .ok_or_else(|| anyhow!("Stored MCP OAuth token does not include a refresh token"))?;
292 let mut form = vec![
293 ("grant_type".to_string(), "refresh_token".to_string()),
294 ("client_id".to_string(), config.client_id.clone()),
295 ("refresh_token".to_string(), refresh_token.to_string()),
296 ];
297 if let Some(audience) = config.audience.as_deref()
298 && !audience.trim().is_empty()
299 {
300 form.push(("audience".to_string(), audience.to_string()));
301 }
302 form.extend(
303 config
304 .extra_token_params
305 .iter()
306 .map(|(key, value)| (key.clone(), value.clone())),
307 );
308
309 let refreshed = send_token_request(&config.token_url, &form).await?;
310 Ok(McpOAuthToken {
311 refresh_token: refreshed.refresh_token.or_else(|| current.refresh_token.clone()),
312 ..refreshed
313 })
314}
315
316async fn send_token_request(token_url: &str, form: &[(String, String)]) -> Result<McpOAuthToken> {
317 let response = Client::new()
318 .post(token_url)
319 .header("Content-Type", "application/x-www-form-urlencoded")
320 .form(form)
321 .send()
322 .await
323 .with_context(|| format!("failed to send MCP OAuth request to {token_url}"))?;
324 let status = response.status();
325 let body = response.text().await.context("failed to read MCP OAuth response body")?;
326
327 if !status.is_success() {
328 bail!("MCP OAuth request failed (HTTP {status}): {body}");
329 }
330
331 let payload: TokenResponse = serde_json::from_str(&body).context("failed to parse MCP OAuth token response")?;
332 let now = now_secs();
333 Ok(McpOAuthToken {
334 access_token: payload.access_token,
335 refresh_token: payload.refresh_token,
336 token_type: payload.token_type,
337 scope: payload.scope,
338 obtained_at: now,
339 expires_at: payload.expires_in.map(|secs| now.saturating_add(secs)),
340 })
341}
342
343#[derive(Debug, Deserialize)]
344struct TokenResponse {
345 access_token: String,
346 #[serde(default)]
347 refresh_token: Option<String>,
348 #[serde(default)]
349 token_type: Option<String>,
350 #[serde(default)]
351 scope: Option<String>,
352 #[serde(default)]
353 expires_in: Option<u64>,
354}
355
356fn generate_state() -> Result<String> {
357 let mut state_bytes = [0_u8; 32];
358 SystemRandom::new()
359 .fill(&mut state_bytes)
360 .map_err(|_| anyhow!("failed to generate MCP OAuth state"))?;
361 Ok(URL_SAFE_NO_PAD.encode(state_bytes))
362}
363
364fn save_token(provider_name: &str, token: &McpOAuthToken, storage_mode: AuthCredentialsStoreMode) -> Result<()> {
365 let serialized = serde_json::to_string(token).context("failed to serialize MCP OAuth token")?;
366 token_storage(provider_name).store_with_mode(&serialized, storage_mode)
367}
368
369fn load_token(provider_name: &str, storage_mode: AuthCredentialsStoreMode) -> Result<Option<McpOAuthToken>> {
370 let Some(serialized) = token_storage(provider_name).load_with_mode(storage_mode)? else {
371 return Ok(None);
372 };
373 serde_json::from_str(&serialized)
374 .context("failed to parse stored MCP OAuth token")
375 .map(Some)
376}
377
378fn clear_token(provider_name: &str, storage_mode: AuthCredentialsStoreMode) -> Result<()> {
379 token_storage(provider_name).clear_with_mode(storage_mode)
380}
381
382fn token_storage(provider_name: &str) -> CredentialStorage {
383 let normalized_provider = provider_name
384 .chars()
385 .map(|ch| {
386 if ch.is_ascii_alphanumeric() || ch == '-' || ch == '_' {
387 ch
388 } else {
389 '_'
390 }
391 })
392 .collect::<String>();
393 CredentialStorage::new("vtcode", format!("mcp_oauth_{normalized_provider}"))
394}
395
396fn now_secs() -> u64 {
397 std::time::SystemTime::now()
398 .duration_since(std::time::UNIX_EPOCH)
399 .map(|duration| duration.as_secs())
400 .unwrap_or(0)
401}
402
403#[cfg(test)]
404mod tests {
405 use super::*;
406 use assert_fs::TempDir;
407 use serial_test::serial;
408 use std::path::PathBuf;
409
410 struct TestAuthDirGuard {
411 previous: Option<PathBuf>,
412 temp_dir: Option<TempDir>,
413 }
414
415 impl TestAuthDirGuard {
416 fn new() -> Self {
417 let temp_dir = TempDir::new().expect("temp dir");
418 let previous =
419 crate::storage_paths::auth_storage_dir_override_for_tests().expect("read previous auth dir override");
420 crate::storage_paths::set_auth_storage_dir_override_for_tests(Some(temp_dir.path().to_path_buf()))
421 .expect("set auth dir override");
422 Self { previous, temp_dir: Some(temp_dir) }
423 }
424 }
425
426 impl Drop for TestAuthDirGuard {
427 fn drop(&mut self) {
428 crate::storage_paths::set_auth_storage_dir_override_for_tests(self.previous.clone())
429 .expect("restore auth dir override");
430 if let Some(temp_dir) = self.temp_dir.take() {
431 drop(temp_dir.close());
432 }
433 }
434 }
435
436 fn sample_config() -> McpOAuthConfig {
437 McpOAuthConfig {
438 authorization_url: "https://example.com/oauth/authorize".to_string(),
439 token_url: "https://example.com/oauth/token".to_string(),
440 client_id: "client-123".to_string(),
441 scopes: vec!["mcp:read".to_string(), "mcp:write".to_string()],
442 audience: Some("mcp-api".to_string()),
443 callback_port: 8123,
444 flow_timeout_secs: 120,
445 credentials_store_mode: AuthCredentialsStoreMode::File,
446 extra_auth_params: BTreeMap::from([("prompt".to_string(), "consent".to_string())]),
447 extra_token_params: BTreeMap::new(),
448 }
449 }
450
451 #[test]
452 fn prepare_login_builds_expected_auth_url() {
453 let service = McpOAuthService::new();
454 let prepared = service.prepare_login("demo", &sample_config()).expect("prepare login");
455
456 assert!(prepared.auth_url.contains("response_type=code"));
457 assert!(prepared.auth_url.contains("client_id=client-123"));
458 assert!(prepared.auth_url.contains("scope=mcp%3Aread+mcp%3Awrite"));
459 assert!(prepared.auth_url.contains("audience=mcp-api"));
460 assert!(prepared.auth_url.contains("prompt=consent"));
461 assert!(prepared.auth_url.contains("code_challenge="));
462 assert!(prepared.auth_url.contains("state="));
463 assert_eq!(prepared.callback_port, 8123);
464 assert_eq!(prepared.timeout_secs, 120);
465 }
466
467 #[test]
468 #[serial]
469 fn status_reflects_stored_token() {
470 let _guard = TestAuthDirGuard::new();
471 let service = McpOAuthService::new();
472 let storage_mode = AuthCredentialsStoreMode::File;
473 assert!(matches!(service.status("demo", storage_mode).expect("status"), McpOAuthStatus::NotAuthenticated));
474
475 save_token(
476 "demo",
477 &McpOAuthToken {
478 access_token: "access".to_string(),
479 refresh_token: Some("refresh".to_string()),
480 token_type: Some("Bearer".to_string()),
481 scope: Some("mcp:read".to_string()),
482 obtained_at: now_secs(),
483 expires_at: Some(now_secs() + 3600),
484 },
485 storage_mode,
486 )
487 .expect("save token");
488
489 let status = service.status("demo", storage_mode).expect("status");
490 assert!(matches!(status, McpOAuthStatus::Authenticated { expires_in: Some(_), .. }));
491 }
492
493 #[test]
494 #[serial]
495 fn logout_clears_stored_token() {
496 let _guard = TestAuthDirGuard::new();
497 let service = McpOAuthService::new();
498 let storage_mode = AuthCredentialsStoreMode::File;
499 save_token(
500 "demo",
501 &McpOAuthToken {
502 access_token: "access".to_string(),
503 refresh_token: None,
504 token_type: Some("Bearer".to_string()),
505 scope: None,
506 obtained_at: now_secs(),
507 expires_at: None,
508 },
509 storage_mode,
510 )
511 .expect("save token");
512
513 drop(service.logout("demo", storage_mode).expect("logout"));
514 assert!(load_token("demo", storage_mode).expect("load").is_none());
515 }
516}