apollo/providers/
oauth.rs1use 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#[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 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 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 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; 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
101pub 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 assert!(!cache.token.blocking_read().is_none());
142 }
143
144 #[test]
145 fn test_load_oauth_fails_gracefully() {
146 let result = load_oauth_token_from_file();
148 let _ = result;
150 }
151}