Skip to main content

claude_codex/providers/kimi/auth/
manager.rs

1use std::sync::{Arc, Mutex};
2use std::time::{SystemTime, UNIX_EPOCH};
3
4use super::constants::{CLIENT_ID, REFRESH_MARGIN_MS, oauth_host};
5use super::headers::common_headers;
6use super::jwt::extract_user_id;
7use super::login::TokenResponse;
8use super::token_store::{KimiTokenStore, StoredAuth};
9use crate::auth::AuthStorage;
10
11const MAX_REFRESH_ATTEMPTS: u32 = 3;
12const RETRYABLE_STATUSES: &[u16] = &[429, 500, 502, 503, 504];
13
14pub struct KimiAuthManager<S: AuthStorage<StoredAuth>> {
15    pub store: KimiTokenStore<S>,
16    cached: Arc<Mutex<Option<StoredAuth>>>,
17}
18
19impl<S: AuthStorage<StoredAuth>> KimiAuthManager<S> {
20    pub fn new(store: KimiTokenStore<S>) -> Self {
21        Self {
22            store,
23            cached: Arc::new(Mutex::new(None)),
24        }
25    }
26
27    fn now_ms() -> u64 {
28        SystemTime::now()
29            .duration_since(UNIX_EPOCH)
30            .unwrap_or_default()
31            .as_millis() as u64
32    }
33
34    pub fn get_auth(&self) -> Result<StoredAuth, anyhow::Error> {
35        let cached = {
36            let guard = self.cached.lock().map_err(|e| anyhow::anyhow!("{e}"))?;
37            guard.clone()
38        };
39        let stored = match cached {
40            Some(ref auth) => auth.clone(),
41            None => {
42                let loaded = self.store.load_auth()?;
43                match loaded {
44                    Some(auth) => {
45                        let mut guard = self.cached.lock().map_err(|e| anyhow::anyhow!("{e}"))?;
46                        *guard = Some(auth.clone());
47                        auth
48                    }
49                    None => {
50                        anyhow::bail!("Not authenticated. Run: claude-codex kimi auth login");
51                    }
52                }
53            }
54        };
55
56        if stored.expires > Self::now_ms() + REFRESH_MARGIN_MS {
57            return Ok(stored);
58        }
59
60        self.refresh_now(&stored)
61    }
62
63    pub fn force_refresh(&self) -> Result<StoredAuth, anyhow::Error> {
64        let stored = {
65            let guard = self.cached.lock().map_err(|e| anyhow::anyhow!("{e}"))?;
66            guard.clone()
67        };
68        let stored = match stored {
69            Some(auth) => auth,
70            None => {
71                let loaded = self.store.load_auth()?;
72                loaded.ok_or_else(|| anyhow::anyhow!("Not authenticated"))?
73            }
74        };
75        self.refresh_now(&stored)
76    }
77
78    fn refresh_now(&self, current: &StoredAuth) -> Result<StoredAuth, anyhow::Error> {
79        if current.refresh.is_empty() {
80            anyhow::bail!("No refresh token stored; re-authenticate");
81        }
82
83        let headers = common_headers()?;
84        let client = reqwest::blocking::Client::new();
85
86        for attempt in 0..MAX_REFRESH_ATTEMPTS {
87            let form = [
88                ("client_id", CLIENT_ID.to_string()),
89                ("grant_type", "refresh_token".to_string()),
90                ("refresh_token", current.refresh.clone()),
91            ];
92
93            let resp = match client
94                .post(format!("{}/api/oauth/token", oauth_host()))
95                .headers(build_headers_map(&headers))
96                .form(&form)
97                .send()
98            {
99                Ok(r) => r,
100                Err(err) => {
101                    if attempt < MAX_REFRESH_ATTEMPTS - 1 {
102                        let ms = 2u64.pow(attempt) * 1000;
103                        std::thread::sleep(std::time::Duration::from_millis(ms));
104                        continue;
105                    }
106                    anyhow::bail!("refresh network error: {err}");
107                }
108            };
109
110            if resp.status().as_u16() == 200 {
111                let tokens: TokenResponse = resp.json()?;
112                let expires = Self::now_ms() + (tokens.expires_in.unwrap_or(900) as u64 * 1000);
113                let next = StoredAuth {
114                    access: tokens.access_token.clone(),
115                    refresh: tokens
116                        .refresh_token
117                        .unwrap_or_else(|| current.refresh.clone()),
118                    expires,
119                    scope: tokens.scope.clone(),
120                    user_id: extract_user_id(&tokens.access_token)
121                        .or_else(|| current.user_id.clone()),
122                };
123                self.store.save_auth(next.clone())?;
124                {
125                    let mut guard = self.cached.lock().map_err(|e| anyhow::anyhow!("{e}"))?;
126                    *guard = Some(next.clone());
127                }
128                return Ok(next);
129            }
130
131            let status = resp.status().as_u16();
132            if status == 401 || status == 403 {
133                {
134                    let mut guard = self.cached.lock().map_err(|e| anyhow::anyhow!("{e}"))?;
135                    *guard = None;
136                }
137                let _ = self.store.clear_auth();
138                let err_msg = resp
139                    .text()
140                    .unwrap_or_else(|_| "Token refresh unauthorized".to_string());
141                anyhow::bail!("{err_msg}");
142            }
143
144            if !RETRYABLE_STATUSES.contains(&status) {
145                anyhow::bail!("Token refresh failed: {status}");
146            }
147
148            if attempt < MAX_REFRESH_ATTEMPTS - 1 {
149                let ms = 2u64.pow(attempt) * 1000;
150                std::thread::sleep(std::time::Duration::from_millis(ms));
151            }
152        }
153
154        anyhow::bail!("Token refresh failed after {MAX_REFRESH_ATTEMPTS} attempts");
155    }
156
157    pub fn persist_initial_tokens(
158        &self,
159        tokens: &TokenResponse,
160    ) -> Result<StoredAuth, anyhow::Error> {
161        let expires = Self::now_ms() + (tokens.expires_in.unwrap_or(900) * 1000);
162        let auth = StoredAuth {
163            access: tokens.access_token.clone(),
164            refresh: tokens.refresh_token.clone().unwrap_or_default(),
165            expires,
166            scope: tokens.scope.clone(),
167            user_id: extract_user_id(&tokens.access_token),
168        };
169        self.store.save_auth(auth.clone())?;
170        {
171            let mut guard = self.cached.lock().map_err(|e| anyhow::anyhow!("{e}"))?;
172            *guard = Some(auth.clone());
173        }
174        Ok(auth)
175    }
176
177    pub fn reset_cache(&self) {
178        if let Ok(mut guard) = self.cached.lock() {
179            *guard = None;
180        }
181    }
182}
183
184fn build_headers_map(
185    headers: &std::collections::HashMap<String, String>,
186) -> reqwest::header::HeaderMap {
187    let mut map = reqwest::header::HeaderMap::new();
188    for (k, v) in headers {
189        if let Ok(name) = reqwest::header::HeaderName::from_bytes(k.as_bytes())
190            && let Ok(value) = reqwest::header::HeaderValue::from_str(v)
191        {
192            map.insert(name, value);
193        }
194    }
195    map
196}
197
198#[cfg(test)]
199mod tests {
200    use super::*;
201    use crate::auth::InMemoryAuthStore;
202
203    fn test_store() -> KimiTokenStore<InMemoryAuthStore<StoredAuth>> {
204        KimiTokenStore::new(InMemoryAuthStore::new())
205    }
206
207    #[test]
208    fn get_auth_returns_stored() {
209        let store = test_store();
210        let auth = StoredAuth {
211            access: "test_access".into(),
212            refresh: "test_refresh".into(),
213            expires: 9999999999999,
214            scope: Some("openid".into()),
215            user_id: Some("user1".into()),
216        };
217        store.save_auth(auth.clone()).unwrap();
218        let manager = KimiAuthManager::new(store);
219        let result = manager.get_auth().unwrap();
220        assert_eq!(result.access, "test_access");
221        assert_eq!(result.user_id.as_deref(), Some("user1"));
222    }
223
224    #[test]
225    fn get_auth_fails_when_no_auth() {
226        let store = test_store();
227        let manager = KimiAuthManager::new(store);
228        assert!(manager.get_auth().is_err());
229    }
230}