claude_codex/providers/kimi/auth/
manager.rs1use 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}