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    audience: Option<String>,
35    /// Local callback server port.
36    pub callback_port: u16,
37    /// Browser-flow timeout in seconds.
38    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    extra_auth_params: BTreeMap<String, String>,
45    /// Extra form fields appended to token exchanges and refreshes.
46    #[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/// Stored OAuth token for an MCP HTTP provider.
87#[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/// 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    success: bool,
133    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
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}