use std::collections::HashMap;
use std::sync::Mutex;
use serde::Deserialize;
#[path = "refresh_state.rs"]
mod refresh_state;
use refresh_state::RefreshAttempts;
use crate::subscription::{SubscriptionProvider, SubscriptionToken};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum BodyStyle {
Json,
Form,
}
#[derive(Debug, Clone, Copy)]
struct RefreshConfig {
token_url: &'static str,
client_id: &'static str,
client_secret_env: Option<&'static str>,
style: BodyStyle,
}
pub const GEMINI_CLIENT_SECRET_ENV: &str = "GEMINI_OAUTH_CLIENT_SECRET";
pub const CLAUDE_CLIENT_ID: &str = "9d1c250a-e61b-44d9-88ed-5944d1962f5e";
pub const CLAUDE_TOKEN_URL: &str = "https://platform.claude.com/v1/oauth/token";
const REFRESH_SKEW_MS: i64 = 60_000;
const fn refresh_config(provider: SubscriptionProvider) -> RefreshConfig {
match provider {
SubscriptionProvider::Claude => RefreshConfig {
token_url: CLAUDE_TOKEN_URL,
client_id: CLAUDE_CLIENT_ID,
client_secret_env: None,
style: BodyStyle::Json,
},
SubscriptionProvider::Codex => RefreshConfig {
token_url: "https://auth.openai.com/oauth/token",
client_id: "app_EMoamEEZ73f0CkXaXp7hrann",
client_secret_env: None,
style: BodyStyle::Json,
},
SubscriptionProvider::Gemini => RefreshConfig {
token_url: "https://oauth2.googleapis.com/token",
client_id: "681255809395-oo8ft2oprdrnp9e3aqf6av3hmdib135j.apps.googleusercontent.com",
client_secret_env: Some(GEMINI_CLIENT_SECRET_ENV),
style: BodyStyle::Form,
},
SubscriptionProvider::Qwen => RefreshConfig {
token_url: "https://chat.qwen.ai/api/v1/oauth2/token",
client_id: "f0304373b74a44d2b584a3fb70ca9e56",
client_secret_env: None,
style: BodyStyle::Form,
},
}
}
fn encode_form(pairs: &[(&str, &str)]) -> String {
fn encode(value: &str) -> String {
let mut out = String::with_capacity(value.len());
for byte in value.bytes() {
if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b'~') {
out.push(byte as char);
} else {
out.push('%');
out.push(
char::from_digit(u32::from(byte >> 4), 16)
.unwrap()
.to_ascii_uppercase(),
);
out.push(
char::from_digit(u32::from(byte & 0x0f), 16)
.unwrap()
.to_ascii_uppercase(),
);
}
}
out
}
pairs
.iter()
.map(|(k, v)| format!("{}={}", encode(k), encode(v)))
.collect::<Vec<_>>()
.join("&")
}
#[derive(Debug, Deserialize, Default)]
struct RefreshResponse {
access_token: Option<String>,
refresh_token: Option<String>,
expires_in: Option<i64>,
}
#[derive(Debug)]
pub enum RefreshError {
Unsupported,
NoRefreshToken,
Request(String),
Status(u16, String, Option<i64>),
Parse(String),
}
const TERMINAL_OAUTH_ERRORS: [&str; 4] = [
"invalid_grant",
"invalid_client",
"unauthorized_client",
"unsupported_grant_type",
];
#[must_use]
fn oauth_error_code(body: &str) -> Option<String> {
let parsed: serde_json::Value = serde_json::from_str(body).ok()?;
let error = parsed.get("error")?;
if let Some(code) = error.as_str() {
return Some(code.to_string());
}
error
.get("type")
.or_else(|| error.get("code"))
.and_then(serde_json::Value::as_str)
.map(str::to_string)
}
impl RefreshError {
#[must_use]
pub fn is_invalid_grant(&self) -> bool {
let Self::Status(code, body, _) = self else {
return false;
};
if !matches!(code, 400 | 401 | 403) {
return false;
}
oauth_error_code(body).is_some_and(|code| TERMINAL_OAUTH_ERRORS.contains(&code.as_str()))
}
#[must_use]
pub const fn is_rate_limited(&self) -> bool {
matches!(self, Self::Status(429, _, _))
}
#[must_use]
pub const fn retry_after_ms(&self) -> Option<i64> {
match self {
Self::Status(_, _, Some(seconds)) => Some(seconds.saturating_mul(1_000)),
_ => None,
}
}
}
impl std::fmt::Display for RefreshError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Unsupported => write!(f, "provider does not support router-driven refresh"),
Self::NoRefreshToken => write!(f, "no refresh token available"),
Self::Request(m) => write!(f, "refresh request failed: {m}"),
Self::Status(_, m, _) if self.is_invalid_grant() => write!(
f,
"refresh token is no longer valid (invalid_grant) — re-authenticate this \
subscription with `link-assistant-router auth <provider>`; waiting will not \
help: {m}"
),
Self::Status(429, m, _) => write!(
f,
"refresh endpoint rate-limited this request (429); it will be retried \
automatically and the subscription remains usable: {m}"
),
Self::Status(code, m, _) => write!(f, "refresh endpoint returned {code}: {m}"),
Self::Parse(m) => write!(f, "refresh response parse error: {m}"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CredentialEvidence {
Working,
Rejected,
}
impl std::error::Error for RefreshError {}
fn merge_refresh_response(
prev: &SubscriptionToken,
resp: &RefreshResponse,
now_ms: i64,
) -> Option<SubscriptionToken> {
let access_token = resp.access_token.clone().filter(|s| !s.is_empty())?;
Some(SubscriptionToken {
access_token,
refresh_token: resp
.refresh_token
.clone()
.filter(|s| !s.is_empty())
.or_else(|| prev.refresh_token.clone()),
expires_at_ms: resp.expires_in.map(|secs| now_ms + secs * 1000),
account_id: prev.account_id.clone(),
resource_url: prev.resource_url.clone(),
})
}
pub async fn refresh(
client: &reqwest::Client,
provider: SubscriptionProvider,
prev: &SubscriptionToken,
now_ms: i64,
) -> Result<SubscriptionToken, RefreshError> {
refresh_at(
client,
refresh_config(provider).token_url,
provider,
prev,
now_ms,
)
.await
}
async fn refresh_at(
client: &reqwest::Client,
token_url: &str,
provider: SubscriptionProvider,
prev: &SubscriptionToken,
now_ms: i64,
) -> Result<SubscriptionToken, RefreshError> {
let config = refresh_config(provider);
let refresh_token = prev
.refresh_token
.as_deref()
.filter(|s| !s.is_empty())
.ok_or(RefreshError::NoRefreshToken)?;
let client_secret = config
.client_secret_env
.and_then(|key| std::env::var(key).ok())
.filter(|s| !s.is_empty());
let request = match config.style {
BodyStyle::Json => {
let mut body = serde_json::json!({
"grant_type": "refresh_token",
"refresh_token": refresh_token,
"client_id": config.client_id,
});
if let Some(secret) = client_secret.as_deref() {
body["client_secret"] = serde_json::Value::String(secret.to_string());
}
client.post(token_url).json(&body)
}
BodyStyle::Form => {
let mut form = vec![
("grant_type", "refresh_token"),
("refresh_token", refresh_token),
("client_id", config.client_id),
];
if let Some(secret) = client_secret.as_deref() {
form.push(("client_secret", secret));
}
client
.post(token_url)
.header("content-type", "application/x-www-form-urlencoded")
.body(encode_form(&form))
}
};
let response = request
.send()
.await
.map_err(|e| RefreshError::Request(e.to_string()))?;
let status = response.status();
if !status.is_success() {
let retry_after = crate::request_routing::retry_after_duration(response.headers())
.and_then(|delay| i64::try_from(delay.as_secs()).ok());
let body = response.text().await.unwrap_or_default();
return Err(RefreshError::Status(status.as_u16(), body, retry_after));
}
let parsed: RefreshResponse = response
.json()
.await
.map_err(|e| RefreshError::Parse(e.to_string()))?;
merge_refresh_response(prev, &parsed, now_ms)
.ok_or_else(|| RefreshError::Parse("response contained no access_token".to_string()))
}
#[derive(Debug, Default)]
pub struct TokenCache {
inner: Mutex<HashMap<(SubscriptionProvider, String), SubscriptionToken>>,
attempts: RefreshAttempts,
evidence: Mutex<HashMap<SubscriptionProvider, CredentialEvidence>>,
refresh_errors: Mutex<HashMap<SubscriptionProvider, String>>,
}
impl TokenCache {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub async fn get_fresh(
&self,
client: &reqwest::Client,
provider: SubscriptionProvider,
disk_token: SubscriptionToken,
now_ms: i64,
) -> SubscriptionToken {
self.get_fresh_for(client, provider, "primary", disk_token, now_ms)
.await
}
pub async fn get_fresh_for(
&self,
client: &reqwest::Client,
provider: SubscriptionProvider,
account: &str,
disk_token: SubscriptionToken,
now_ms: i64,
) -> SubscriptionToken {
self.get_fresh_for_at(
client,
refresh_config(provider).token_url,
provider,
account,
disk_token,
now_ms,
)
.await
}
pub async fn refresh_rejected(
&self,
client: &reqwest::Client,
provider: SubscriptionProvider,
account: &str,
disk_token: SubscriptionToken,
now_ms: i64,
) -> Option<SubscriptionToken> {
self.refresh_rejected_at(
client,
refresh_config(provider).token_url,
provider,
account,
disk_token,
now_ms,
)
.await
}
async fn refresh_rejected_at(
&self,
client: &reqwest::Client,
token_url: &str,
provider: SubscriptionProvider,
account: &str,
rejected: SubscriptionToken,
now_ms: i64,
) -> Option<SubscriptionToken> {
let attempt = self.attempts.for_subscription(provider, account, &rejected);
let mut attempt = attempt.lock().await;
if let Some(cached) = self.cached_valid_for(provider, account, now_ms)
&& cached.access_token != rejected.access_token
{
return Some(cached);
}
if attempt.suppresses_attempt(now_ms) {
return None;
}
match refresh_at(client, token_url, provider, &rejected, now_ms).await {
Ok(fresh) => {
self.store_for(provider, account, fresh.clone());
attempt.record_success();
self.record_credential_working(provider);
if let Ok(mut guard) = self.refresh_errors.lock() {
guard.remove(&provider);
}
(fresh.access_token != rejected.access_token).then_some(fresh)
}
Err(error) => {
tracing::warn!("refresh after a rejected {provider} token failed: {error}");
self.record_refresh_error(provider, &error.to_string());
if error.is_invalid_grant() {
attempt.record_terminal_failure();
self.record_credential_rejected(provider);
} else {
attempt.record_transient_failure_after(now_ms, error.retry_after_ms());
}
None
}
}
}
async fn get_fresh_for_at(
&self,
client: &reqwest::Client,
token_url: &str,
provider: SubscriptionProvider,
account: &str,
disk_token: SubscriptionToken,
now_ms: i64,
) -> SubscriptionToken {
let key = (provider, account.to_string());
let attempt = self
.attempts
.for_subscription(provider, account, &disk_token);
let mut attempt = attempt.lock().await;
if attempt.reset_if_changed(&disk_token) {
if let Ok(mut guard) = self.inner.lock() {
guard.remove(&key);
}
if let Ok(mut guard) = self.evidence.lock() {
guard.remove(&provider);
}
}
if !disk_token.is_expired(now_ms.saturating_add(REFRESH_SKEW_MS)) {
return disk_token;
}
if let Some(cached) = self.cached_valid_for(provider, account, now_ms) {
return cached;
}
if attempt.suppresses_attempt(now_ms) {
return disk_token;
}
match refresh_at(client, token_url, provider, &disk_token, now_ms).await {
Ok(fresh) => {
self.store_for(provider, account, fresh.clone());
attempt.record_success();
self.record_credential_working(provider);
if let Ok(mut guard) = self.refresh_errors.lock() {
guard.remove(&provider);
}
fresh
}
Err(e) => {
tracing::warn!("subscription token refresh for {provider} failed: {e}");
self.record_refresh_error(provider, &e.to_string());
if e.is_invalid_grant() {
attempt.record_terminal_failure();
self.record_credential_rejected(provider);
} else {
attempt.record_transient_failure_after(now_ms, e.retry_after_ms());
}
disk_token
}
}
}
pub fn record_credential_working(&self, provider: SubscriptionProvider) {
self.record_evidence(provider, CredentialEvidence::Working);
}
pub fn record_status(&self, provider: SubscriptionProvider, status: u16) {
if status == 401 || status == 403 {
self.record_credential_rejected(provider);
} else if (200..300).contains(&status) {
self.record_credential_working(provider);
}
}
pub fn record_credential_rejected(&self, provider: SubscriptionProvider) {
self.record_evidence(provider, CredentialEvidence::Rejected);
}
#[must_use]
pub fn evidence(&self, provider: SubscriptionProvider) -> Option<CredentialEvidence> {
self.evidence
.lock()
.ok()
.and_then(|guard| guard.get(&provider).copied())
}
#[must_use]
pub fn last_refresh_error(&self, provider: SubscriptionProvider) -> Option<String> {
self.refresh_errors
.lock()
.ok()
.and_then(|guard| guard.get(&provider).cloned())
}
fn record_evidence(&self, provider: SubscriptionProvider, evidence: CredentialEvidence) {
if let Ok(mut guard) = self.evidence.lock() {
guard.insert(provider, evidence);
}
}
fn record_refresh_error(&self, provider: SubscriptionProvider, error: &str) {
if let Ok(mut guard) = self.refresh_errors.lock() {
guard.insert(provider, error.to_string());
}
}
pub fn store_refreshed(
&self,
provider: SubscriptionProvider,
account: &str,
token: SubscriptionToken,
) {
self.store_for(provider, account, token);
}
fn cached_valid_for(
&self,
provider: SubscriptionProvider,
account: &str,
now_ms: i64,
) -> Option<SubscriptionToken> {
let guard = self.inner.lock().ok()?;
guard
.get(&(provider, account.to_string()))
.filter(|token| !token.is_expired(now_ms))
.cloned()
}
fn store_for(&self, provider: SubscriptionProvider, account: &str, token: SubscriptionToken) {
if let Ok(mut guard) = self.inner.lock() {
guard.insert((provider, account.to_string()), token);
}
}
}
#[cfg(test)]
#[path = "refresh_tests.rs"]
mod tests;