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 pub audience: Option<String>,
35 pub callback_port: u16,
37 pub flow_timeout_secs: u64,
39 #[serde(default)]
41 pub credentials_store_mode: AuthCredentialsStoreMode,
42 #[serde(default)]
44 pub extra_auth_params: BTreeMap<String, String>,
45 #[serde(default)]
47 pub 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 pub 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 pub access_token: String,
90 pub refresh_token: Option<String>,
91 pub token_type: Option<String>,
92 pub scope: Option<String>,
93 pub obtained_at: u64,
94 pub expires_at: Option<u64>,
95}
96
97impl McpOAuthToken {
98 pub 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 pub success: bool,
133 pub 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
229fn build_auth_url(config: &McpOAuthConfig, challenge: &PkceChallenge, state: &str) -> Result<String> {
230 let mut url = Url::parse(&config.authorization_url).context("invalid oauth.authorization_url")?;
231 {
232 let mut query = url.query_pairs_mut();
233 query.append_pair("response_type", "code");
234 query.append_pair("client_id", &config.client_id);
235 query.append_pair("redirect_uri", &config.callback_url());
236 query.append_pair("code_challenge", &challenge.code_challenge);
237 query.append_pair("code_challenge_method", &challenge.code_challenge_method);
238 query.append_pair("state", state);
239 if !config.scopes.is_empty() {
240 query.append_pair("scope", &config.scopes.join(" "));
241 }
242 if let Some(audience) = config.audience.as_deref()
243 && !audience.trim().is_empty()
244 {
245 query.append_pair("audience", audience);
246 }
247 for (key, value) in &config.extra_auth_params {
248 if !key.trim().is_empty() {
249 query.append_pair(key, value);
250 }
251 }
252 }
253 Ok(url.to_string())
254}
255
256async fn exchange_code_for_token(
257 config: &McpOAuthConfig,
258 code: &str,
259 challenge: &PkceChallenge,
260) -> Result<McpOAuthToken> {
261 let mut form = vec![
262 ("grant_type".to_string(), "authorization_code".to_string()),
263 ("client_id".to_string(), config.client_id.clone()),
264 ("code".to_string(), code.to_string()),
265 ("redirect_uri".to_string(), config.callback_url()),
266 ("code_verifier".to_string(), challenge.code_verifier.to_string()),
267 ];
268 if let Some(audience) = config.audience.as_deref()
269 && !audience.trim().is_empty()
270 {
271 form.push(("audience".to_string(), audience.to_string()));
272 }
273 form.extend(
274 config
275 .extra_token_params
276 .iter()
277 .map(|(key, value)| (key.clone(), value.clone())),
278 );
279 send_token_request(&config.token_url, &form).await
280}
281
282async fn refresh_token(config: &McpOAuthConfig, current: &McpOAuthToken) -> Result<McpOAuthToken> {
283 let refresh_token = current
284 .refresh_token
285 .as_deref()
286 .filter(|value| !value.trim().is_empty())
287 .ok_or_else(|| anyhow!("Stored MCP OAuth token does not include a refresh token"))?;
288 let mut form = vec![
289 ("grant_type".to_string(), "refresh_token".to_string()),
290 ("client_id".to_string(), config.client_id.clone()),
291 ("refresh_token".to_string(), refresh_token.to_string()),
292 ];
293 if let Some(audience) = config.audience.as_deref()
294 && !audience.trim().is_empty()
295 {
296 form.push(("audience".to_string(), audience.to_string()));
297 }
298 form.extend(
299 config
300 .extra_token_params
301 .iter()
302 .map(|(key, value)| (key.clone(), value.clone())),
303 );
304
305 let refreshed = send_token_request(&config.token_url, &form).await?;
306 Ok(McpOAuthToken {
307 refresh_token: refreshed.refresh_token.or_else(|| current.refresh_token.clone()),
308 ..refreshed
309 })
310}
311
312async fn send_token_request(token_url: &str, form: &[(String, String)]) -> Result<McpOAuthToken> {
313 let response = Client::new()
314 .post(token_url)
315 .header("Content-Type", "application/x-www-form-urlencoded")
316 .form(form)
317 .send()
318 .await
319 .with_context(|| format!("failed to send MCP OAuth request to {token_url}"))?;
320 let status = response.status();
321 let body = response.text().await.context("failed to read MCP OAuth response body")?;
322
323 if !status.is_success() {
324 bail!("MCP OAuth request failed (HTTP {status}): {body}");
325 }
326
327 let payload: TokenResponse = serde_json::from_str(&body).context("failed to parse MCP OAuth token response")?;
328 let now = now_secs();
329 Ok(McpOAuthToken {
330 access_token: payload.access_token,
331 refresh_token: payload.refresh_token,
332 token_type: payload.token_type,
333 scope: payload.scope,
334 obtained_at: now,
335 expires_at: payload.expires_in.map(|secs| now.saturating_add(secs)),
336 })
337}
338
339#[derive(Debug, Deserialize)]
340struct TokenResponse {
341 access_token: String,
342 #[serde(default)]
343 refresh_token: Option<String>,
344 #[serde(default)]
345 token_type: Option<String>,
346 #[serde(default)]
347 scope: Option<String>,
348 #[serde(default)]
349 expires_in: Option<u64>,
350}
351
352fn generate_state() -> Result<String> {
353 let mut state_bytes = [0_u8; 32];
354 SystemRandom::new()
355 .fill(&mut state_bytes)
356 .map_err(|_| anyhow!("failed to generate MCP OAuth state"))?;
357 Ok(URL_SAFE_NO_PAD.encode(state_bytes))
358}
359
360fn save_token(provider_name: &str, token: &McpOAuthToken, storage_mode: AuthCredentialsStoreMode) -> Result<()> {
361 let serialized = serde_json::to_string(token).context("failed to serialize MCP OAuth token")?;
362 token_storage(provider_name).store_with_mode(&serialized, storage_mode)
363}
364
365fn load_token(provider_name: &str, storage_mode: AuthCredentialsStoreMode) -> Result<Option<McpOAuthToken>> {
366 let Some(serialized) = token_storage(provider_name).load_with_mode(storage_mode)? else {
367 return Ok(None);
368 };
369 serde_json::from_str(&serialized)
370 .context("failed to parse stored MCP OAuth token")
371 .map(Some)
372}
373
374fn clear_token(provider_name: &str, storage_mode: AuthCredentialsStoreMode) -> Result<()> {
375 token_storage(provider_name).clear_with_mode(storage_mode)
376}
377
378fn token_storage(provider_name: &str) -> CredentialStorage {
379 let normalized_provider = provider_name
380 .chars()
381 .map(|ch| {
382 if ch.is_ascii_alphanumeric() || ch == '-' || ch == '_' {
383 ch
384 } else {
385 '_'
386 }
387 })
388 .collect::<String>();
389 CredentialStorage::new("vtcode", format!("mcp_oauth_{normalized_provider}"))
390}
391
392fn now_secs() -> u64 {
393 std::time::SystemTime::now()
394 .duration_since(std::time::UNIX_EPOCH)
395 .map(|duration| duration.as_secs())
396 .unwrap_or(0)
397}
398
399#[cfg(test)]
400mod tests {
401 use super::*;
402 use assert_fs::TempDir;
403 use serial_test::serial;
404 use std::path::PathBuf;
405
406 struct TestAuthDirGuard {
407 previous: Option<PathBuf>,
408 temp_dir: Option<TempDir>,
409 }
410
411 impl TestAuthDirGuard {
412 fn new() -> Self {
413 let temp_dir = TempDir::new().expect("temp dir");
414 let previous =
415 crate::storage_paths::auth_storage_dir_override_for_tests().expect("read previous auth dir override");
416 crate::storage_paths::set_auth_storage_dir_override_for_tests(Some(temp_dir.path().to_path_buf()))
417 .expect("set auth dir override");
418 Self { previous, temp_dir: Some(temp_dir) }
419 }
420 }
421
422 impl Drop for TestAuthDirGuard {
423 fn drop(&mut self) {
424 crate::storage_paths::set_auth_storage_dir_override_for_tests(self.previous.clone())
425 .expect("restore auth dir override");
426 if let Some(temp_dir) = self.temp_dir.take() {
427 let _ = temp_dir.close();
428 }
429 }
430 }
431
432 fn sample_config() -> McpOAuthConfig {
433 McpOAuthConfig {
434 authorization_url: "https://example.com/oauth/authorize".to_string(),
435 token_url: "https://example.com/oauth/token".to_string(),
436 client_id: "client-123".to_string(),
437 scopes: vec!["mcp:read".to_string(), "mcp:write".to_string()],
438 audience: Some("mcp-api".to_string()),
439 callback_port: 8123,
440 flow_timeout_secs: 120,
441 credentials_store_mode: AuthCredentialsStoreMode::File,
442 extra_auth_params: BTreeMap::from([("prompt".to_string(), "consent".to_string())]),
443 extra_token_params: BTreeMap::new(),
444 }
445 }
446
447 #[test]
448 fn prepare_login_builds_expected_auth_url() {
449 let service = McpOAuthService::new();
450 let prepared = service.prepare_login("demo", &sample_config()).expect("prepare login");
451
452 assert!(prepared.auth_url.contains("response_type=code"));
453 assert!(prepared.auth_url.contains("client_id=client-123"));
454 assert!(prepared.auth_url.contains("scope=mcp%3Aread+mcp%3Awrite"));
455 assert!(prepared.auth_url.contains("audience=mcp-api"));
456 assert!(prepared.auth_url.contains("prompt=consent"));
457 assert!(prepared.auth_url.contains("code_challenge="));
458 assert!(prepared.auth_url.contains("state="));
459 assert_eq!(prepared.callback_port, 8123);
460 assert_eq!(prepared.timeout_secs, 120);
461 }
462
463 #[test]
464 #[serial]
465 fn status_reflects_stored_token() {
466 let _guard = TestAuthDirGuard::new();
467 let service = McpOAuthService::new();
468 let storage_mode = AuthCredentialsStoreMode::File;
469 assert!(matches!(service.status("demo", storage_mode).expect("status"), McpOAuthStatus::NotAuthenticated));
470
471 save_token(
472 "demo",
473 &McpOAuthToken {
474 access_token: "access".to_string(),
475 refresh_token: Some("refresh".to_string()),
476 token_type: Some("Bearer".to_string()),
477 scope: Some("mcp:read".to_string()),
478 obtained_at: now_secs(),
479 expires_at: Some(now_secs() + 3600),
480 },
481 storage_mode,
482 )
483 .expect("save token");
484
485 let status = service.status("demo", storage_mode).expect("status");
486 assert!(matches!(status, McpOAuthStatus::Authenticated { expires_in: Some(_), .. }));
487 }
488
489 #[test]
490 #[serial]
491 fn logout_clears_stored_token() {
492 let _guard = TestAuthDirGuard::new();
493 let service = McpOAuthService::new();
494 let storage_mode = AuthCredentialsStoreMode::File;
495 save_token(
496 "demo",
497 &McpOAuthToken {
498 access_token: "access".to_string(),
499 refresh_token: None,
500 token_type: Some("Bearer".to_string()),
501 scope: None,
502 obtained_at: now_secs(),
503 expires_at: None,
504 },
505 storage_mode,
506 )
507 .expect("save token");
508
509 service.logout("demo", storage_mode).expect("logout");
510 assert!(load_token("demo", storage_mode).expect("load").is_none());
511 }
512}