Skip to main content

apollo/providers/
oauth.rs

1//! OAuth token support for Anthropic
2//! Converts Claude.dev OAuth tokens (oat01) to API calls via token exchange
3
4use serde::Deserialize;
5use serde_json::Value;
6use std::sync::Arc;
7use tokio::sync::RwLock;
8
9#[derive(Deserialize)]
10struct RefreshResponse {
11    access_token: String,
12    refresh_token: Option<String>,
13    expires_in: Option<i64>,
14}
15
16/// OAuth token cache (refreshed as needed)
17#[derive(Clone)]
18pub struct OAuthTokenCache {
19    token: Arc<RwLock<Option<String>>>,
20    refresh_token: Arc<RwLock<Option<String>>>,
21    expires_at: Arc<RwLock<i64>>,
22}
23
24impl OAuthTokenCache {
25    pub fn new(initial_token: String, refresh_token: Option<String>, expires_at: i64) -> Self {
26        Self {
27            token: Arc::new(RwLock::new(Some(initial_token))),
28            refresh_token: Arc::new(RwLock::new(refresh_token)),
29            expires_at: Arc::new(RwLock::new(expires_at)),
30        }
31    }
32
33    /// Get current token, refresh if expired
34    pub async fn get_token(&self) -> anyhow::Result<String> {
35        let token = self.token.read().await;
36        if let Some(t) = token.as_ref() {
37            let expires = *self.expires_at.read().await;
38            if expires > chrono::Utc::now().timestamp_millis() {
39                return Ok(t.clone());
40            }
41        }
42        drop(token);
43
44        // Token expired, try refresh
45        self.refresh().await?;
46
47        let token = self.token.read().await;
48        token
49            .clone()
50            .ok_or_else(|| anyhow::anyhow!("Failed to get valid token"))
51    }
52
53    /// Refresh token from Anthropic OAuth endpoint
54    async fn refresh(&self) -> anyhow::Result<()> {
55        let refresh_token = {
56            let rt = self.refresh_token.read().await;
57            rt.clone()
58                .ok_or_else(|| anyhow::anyhow!("No refresh token available"))?
59        };
60
61        let body = Self::fetch_new_token(&refresh_token).await?;
62
63        let expires_in = body.expires_in.unwrap_or(3600) * 1000; // Convert to ms
64        let new_expires = chrono::Utc::now().timestamp_millis() + expires_in;
65
66        *self.token.write().await = Some(body.access_token);
67        if let Some(r) = body.refresh_token {
68            *self.refresh_token.write().await = Some(r);
69        }
70        *self.expires_at.write().await = new_expires;
71
72        Ok(())
73    }
74
75    async fn fetch_new_token(refresh_token: &str) -> anyhow::Result<RefreshResponse> {
76        let client = reqwest::Client::builder()
77            .timeout(std::time::Duration::from_secs(30))
78            .build()?;
79
80        let response = client
81            .post("https://api.anthropic.com/v1/oauth/token")
82            .json(&serde_json::json!({
83                "grant_type": "refresh_token",
84                "refresh_token": refresh_token,
85            }))
86            .send()
87            .await?;
88
89        if !response.status().is_success() {
90            return Err(anyhow::anyhow!(
91                "OAuth token refresh failed: {}",
92                response.status()
93            ));
94        }
95
96        let body = response.json::<RefreshResponse>().await?;
97        Ok(body)
98    }
99}
100
101/// Load OAuth token from Claude.dev credentials file
102pub fn load_oauth_token_from_file() -> anyhow::Result<(String, Option<String>, i64)> {
103    let credentials_path = dirs::home_dir()
104        .ok_or_else(|| anyhow::anyhow!("Cannot determine home directory"))?
105        .join(".claude")
106        .join(".credentials.json");
107
108    let content = std::fs::read_to_string(&credentials_path)
109        .map_err(|e| anyhow::anyhow!("Failed to read Claude credentials: {}", e))?;
110
111    let creds: Value = serde_json::from_str(&content)
112        .map_err(|e| anyhow::anyhow!("Failed to parse Claude credentials: {}", e))?;
113
114    let oauth = &creds["claudeAiOauth"];
115    let access_token = oauth["accessToken"]
116        .as_str()
117        .ok_or_else(|| anyhow::anyhow!("No accessToken in credentials"))?
118        .to_string();
119
120    let refresh_token = oauth["refreshToken"].as_str().map(|s| s.to_string());
121    let expires_at = oauth["expiresAt"]
122        .as_i64()
123        .unwrap_or_else(|| chrono::Utc::now().timestamp_millis() + 3600 * 1000);
124
125    Ok((access_token, refresh_token, expires_at))
126}
127
128#[cfg(test)]
129mod tests {
130    use super::*;
131
132    #[test]
133    fn test_oauth_cache() {
134        let cache = OAuthTokenCache::new(
135            "token123".to_string(),
136            None,
137            chrono::Utc::now().timestamp_millis() + 3600 * 1000,
138        );
139
140        // Should not panic
141        assert!(!cache.token.blocking_read().is_none());
142    }
143
144    #[test]
145    fn test_load_oauth_fails_gracefully() {
146        // This will fail if credentials don't exist, which is expected
147        let result = load_oauth_token_from_file();
148        // Either it succeeds or it fails gracefully
149        let _ = result;
150    }
151}