use anyhow::{Context, Result};
use reqwest::StatusCode;
use reqwest::blocking::RequestBuilder;
use serde_json::Value;
use std::path::Path;
use std::time::SystemTime;
pub const REFRESH_COOLDOWN_SECS: i64 = 300;
pub const EXPIRY_SKEW_SECS: i64 = 60;
#[derive(Default)]
pub struct RefreshCooldown {
until: i64,
mtime: Option<SystemTime>,
}
impl RefreshCooldown {
pub fn active(&self, now: i64, cur_mtime: Option<SystemTime>) -> bool {
now < self.until && self.mtime == cur_mtime
}
pub fn arm(&mut self, now: i64, mtime: Option<SystemTime>) {
self.until = now + REFRESH_COOLDOWN_SECS;
self.mtime = mtime;
}
pub fn clear(&mut self) {
self.until = 0;
self.mtime = None;
}
}
pub fn file_mtime(path: &Path) -> Option<SystemTime> {
std::fs::metadata(path).and_then(|m| m.modified()).ok()
}
pub fn is_expiring(expires_at_secs: Option<i64>, now_secs: i64, skew_secs: i64) -> bool {
expires_at_secs.is_some_and(|exp| exp - now_secs <= skew_secs)
}
pub fn send_refresh(req: RequestBuilder) -> Result<(StatusCode, String)> {
let resp = req.send().context("token refresh request failed")?;
let status = resp.status();
let text = resp.text().unwrap_or_default();
Ok((status, text))
}
pub fn update_json_file_in_place(
path: &Path,
expected_mtime: Option<SystemTime>,
mutate: impl FnOnce(&mut Value) -> Result<()>,
) -> Result<bool> {
if file_mtime(path) != expected_mtime {
return Ok(false);
}
let body = std::fs::read_to_string(path)
.with_context(|| format!("Failed to read {}", path.display()))?;
let mut root: Value = serde_json::from_str(&body).context("Failed to parse credential file")?;
mutate(&mut root)?;
if file_mtime(path) != expected_mtime {
return Ok(false);
}
crate::utils::write_json_atomic_pretty(path, &root)?;
Ok(true)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn is_expiring_unknown_is_never_expiring() {
assert!(!is_expiring(None, 1000, 60));
}
#[test]
fn is_expiring_boundaries() {
assert!(!is_expiring(Some(1061), 1000, 60));
assert!(is_expiring(Some(1060), 1000, 60));
assert!(is_expiring(Some(900), 1000, 60));
}
#[test]
fn cooldown_respects_mtime_change() {
let mt0 = SystemTime::UNIX_EPOCH;
let mt1 = SystemTime::UNIX_EPOCH + std::time::Duration::from_secs(5);
let mut cd = RefreshCooldown::default();
cd.arm(1000, Some(mt0));
assert!(cd.active(1100, Some(mt0)));
assert!(!cd.active(1100, Some(mt1)));
assert!(!cd.active(2000, Some(mt0)));
}
#[test]
fn update_in_place_preserves_unknown_and_aborts_on_mtime_change() {
use std::io::Write;
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("auth.json");
let mut f = std::fs::File::create(&path).unwrap();
write!(
f,
r#"{{"auth_mode":"chatgpt","OPENAI_API_KEY":null,"tokens":{{"access_token":"old","refresh_token":"oldr"}},"extra":123}}"#
)
.unwrap();
drop(f);
let mtime = file_mtime(&path);
let wrote = update_json_file_in_place(&path, mtime, |root| {
let t = root
.get_mut("tokens")
.and_then(|v| v.as_object_mut())
.unwrap();
t.insert("access_token".into(), serde_json::json!("new"));
t.insert("refresh_token".into(), serde_json::json!("newr"));
Ok(())
})
.unwrap();
assert!(wrote);
let v: Value = serde_json::from_str(&std::fs::read_to_string(&path).unwrap()).unwrap();
assert_eq!(v["tokens"]["access_token"], "new");
assert_eq!(v["tokens"]["refresh_token"], "newr");
assert_eq!(v["auth_mode"], "chatgpt");
assert_eq!(v["extra"], 123);
assert!(v.get("OPENAI_API_KEY").is_some());
let wrote2 = update_json_file_in_place(&path, Some(SystemTime::UNIX_EPOCH), |root| {
root["tokens"]["access_token"] = serde_json::json!("should-not-write");
Ok(())
})
.unwrap();
assert!(!wrote2);
let v2: Value = serde_json::from_str(&std::fs::read_to_string(&path).unwrap()).unwrap();
assert_eq!(v2["tokens"]["access_token"], "new");
}
#[test]
fn claude_write_back_preserves_design_oauth_and_ms_expiry() {
use std::io::Write;
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(".credentials.json");
let mut f = std::fs::File::create(&path).unwrap();
write!(
f,
r#"{{"claudeAiOauth":{{"accessToken":"old","refreshToken":"oldr","expiresAt":1,"scopes":["user:inference"],"subscriptionType":"max","rateLimitTier":"t"}},"designOauth":{{"accessToken":"design"}}}}"#
)
.unwrap();
drop(f);
let mtime = file_mtime(&path);
let wrote = update_json_file_in_place(&path, mtime, |root| {
let o = root
.get_mut("claudeAiOauth")
.and_then(|v| v.as_object_mut())
.unwrap();
o.insert("accessToken".into(), serde_json::json!("newacc"));
o.insert("refreshToken".into(), serde_json::json!("newref"));
o.insert("expiresAt".into(), serde_json::json!(1783108188604i64));
Ok(())
})
.unwrap();
assert!(wrote);
let v: Value = serde_json::from_str(&std::fs::read_to_string(&path).unwrap()).unwrap();
assert_eq!(v["claudeAiOauth"]["accessToken"], "newacc");
assert_eq!(v["claudeAiOauth"]["refreshToken"], "newref");
assert_eq!(v["claudeAiOauth"]["expiresAt"], 1783108188604i64);
assert!(v["claudeAiOauth"]["expiresAt"].is_number());
assert_eq!(v["claudeAiOauth"]["subscriptionType"], "max");
assert_eq!(v["claudeAiOauth"]["rateLimitTier"], "t");
assert_eq!(v["designOauth"]["accessToken"], "design");
}
}