Skip to main content

vtcode_auth/
mcp_oauth.rs

1//! OAuth support for HTTP MCP providers.
2
3use 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/// Configuration for OAuth-enabled MCP HTTP providers.
19#[derive(Debug, Clone, Serialize, Deserialize)]
20#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
21#[serde(default)]
22pub struct McpOAuthConfig {
23    /// OAuth authorization endpoint.
24    pub authorization_url: String,
25    /// OAuth token endpoint.
26    pub token_url: String,
27    /// OAuth client identifier.
28    pub client_id: String,
29    /// Requested scopes.
30    #[serde(default)]
31    pub scopes: Vec<String>,
32    /// Optional audience/resource hint sent with the auth and token requests.
33    #[serde(default)]
34    pub audience: Option<String>,
35    /// Local callback server port.
36    pub callback_port: u16,
37    /// Browser-flow timeout in seconds.
38    pub flow_timeout_secs: u64,
39    /// Credential storage backend for this provider's token.
40    #[serde(default)]
41    pub credentials_store_mode: AuthCredentialsStoreMode,
42    /// Extra query parameters appended to the authorization URL.
43    #[serde(default)]
44    pub extra_auth_params: BTreeMap<String, String>,
45    /// Extra form fields appended to token exchanges and refreshes.
46    #[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/// Stored OAuth token for an MCP HTTP provider.
87#[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/// Status for an MCP provider's stored OAuth token.
105#[derive(Debug, Clone)]
106pub enum McpOAuthStatus {
107    Authenticated { age_seconds: u64, expires_in: Option<u64> },
108    NotAuthenticated,
109}
110
111/// Prepared browser-login flow for an MCP OAuth provider.
112#[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/// Completion payload kept intentionally close to Codex app-server.
129#[derive(Debug, Clone, PartialEq, Eq)]
130pub struct McpOAuthLoginCompletion {
131    pub name: String,
132    pub success: bool,
133    pub error: Option<String>,
134}
135
136/// Service for loading, refreshing, and persisting MCP OAuth tokens.
137#[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}