1use serde::{Deserialize, Serialize};
11
12use crate::error::{AppError, Result};
13
14pub const TOKEN_URL: &str = "https://auth.openai.com/oauth/token";
15pub const CLIENT_ID: &str = "app_EMoamEEZ73f0CkXaXp7hrann";
16pub const SCOPE: &str = "openid profile email";
17pub const REFRESH_BUFFER_SECS: i64 = 300;
18
19#[derive(Debug, Serialize)]
20struct RefreshRequest<'a> {
21 client_id: &'a str,
22 grant_type: &'a str,
23 refresh_token: &'a str,
24 scope: &'a str,
25}
26
27#[derive(Debug, Deserialize)]
28pub struct RefreshResponse {
29 #[serde(deserialize_with = "de_nonempty_string")]
30 pub access_token: String,
31 #[serde(default, deserialize_with = "de_opt_nonempty_string")]
32 pub refresh_token: Option<String>,
33 #[serde(default, deserialize_with = "de_opt_nonempty_string")]
34 pub id_token: Option<String>,
35 #[serde(default, deserialize_with = "de_expires_in")]
36 pub expires_in: Option<u64>,
37}
38
39fn de_nonempty_string<'de, D>(d: D) -> std::result::Result<String, D::Error>
40where
41 D: serde::Deserializer<'de>,
42{
43 let value = String::deserialize(d)?;
44 if value.trim().is_empty() {
45 Err(serde::de::Error::custom("token cannot be empty"))
46 } else {
47 Ok(value)
48 }
49}
50
51fn de_opt_nonempty_string<'de, D>(d: D) -> std::result::Result<Option<String>, D::Error>
52where
53 D: serde::Deserializer<'de>,
54{
55 Option::<String>::deserialize(d)?
56 .map(|value| {
57 if value.trim().is_empty() {
58 Err(serde::de::Error::custom("token cannot be empty"))
59 } else {
60 Ok(value)
61 }
62 })
63 .transpose()
64}
65
66fn de_expires_in<'de, D>(d: D) -> std::result::Result<Option<u64>, D::Error>
67where
68 D: serde::Deserializer<'de>,
69{
70 let v = serde_json::Value::deserialize(d)?;
71 match v {
72 serde_json::Value::Null => Ok(None),
73 serde_json::Value::Number(n) => {
74 const MAX_SAFE_EXPIRES_IN: u64 = (i64::MAX as u64) / 2;
75 if let Some(value) = n.as_u64().filter(|value| *value <= MAX_SAFE_EXPIRES_IN) {
76 Ok(Some(value))
77 } else if let Some(value) = n.as_f64()
78 && value.is_finite()
79 && value.fract() == 0.0
80 && (0.0..=MAX_SAFE_EXPIRES_IN as f64).contains(&value)
81 {
82 Ok(Some(value as u64))
83 } else {
84 Err(serde::de::Error::custom(
85 "expires_in must be a non-negative integer in range",
86 ))
87 }
88 }
89 _ => Err(serde::de::Error::custom(
90 "expires_in must be a number or null",
91 )),
92 }
93}
94
95pub async fn refresh(
96 client: &reqwest::Client,
97 endpoint: &str,
98 refresh_token: &str,
99) -> Result<RefreshResponse> {
100 let req = RefreshRequest {
101 client_id: CLIENT_ID,
102 grant_type: "refresh_token",
103 refresh_token,
104 scope: SCOPE,
105 };
106
107 let resp = client
108 .post(endpoint)
109 .header("Content-Type", "application/json")
110 .json(&req)
111 .send()
112 .await?;
113
114 let status = resp.status();
115 let body = crate::vendor::read_body_capped(resp, crate::vendor::MAX_BODY_BYTES).await?;
116 let body = String::from_utf8_lossy(&body).into_owned();
117 if !status.is_success() {
118 let msg = crate::anthropic::oauth::parse_error_body(&body)
119 .unwrap_or_else(|| "Refresh failed".into());
120 return Err(AppError::Http {
121 status: status.as_u16(),
122 body: msg,
123 });
124 }
125 serde_json::from_str(&body).map_err(|e| AppError::Schema(format!("openai token response: {e}")))
126}
127
128pub fn needs_refresh(expires_at_secs: i64, now_secs: i64) -> bool {
129 expires_at_secs < now_secs + REFRESH_BUFFER_SECS
130}
131
132#[cfg(test)]
133mod tests {
134 use super::*;
135
136 #[test]
137 fn needs_refresh_threshold() {
138 let now = 1_000_000;
139 assert!(needs_refresh(now + 100, now));
140 assert!(!needs_refresh(now + 1000, now));
141 }
142
143 #[test]
144 fn malformed_optional_expires_in_is_not_treated_as_absent() {
145 for value in [
146 "3600.5",
147 "-1",
148 "1e300",
149 "18446744073709551615",
150 "true",
151 r#""3600""#,
152 ] {
153 let body = format!(r#"{{"access_token":"new","expires_in":{value}}}"#);
154 assert!(
155 serde_json::from_str::<RefreshResponse>(&body).is_err(),
156 "{body}"
157 );
158 }
159 let response: RefreshResponse =
160 serde_json::from_str(r#"{"access_token":"new","expires_in":null}"#).unwrap();
161 assert_eq!(response.expires_in, None);
162 }
163
164 #[test]
165 fn empty_refresh_tokens_are_schema_drift_not_credentials_to_persist() {
166 for body in [
167 r#"{"access_token":""}"#,
168 r#"{"access_token":"new","refresh_token":" "}"#,
169 r#"{"access_token":"new","id_token":""}"#,
170 ] {
171 assert!(
172 serde_json::from_str::<RefreshResponse>(body).is_err(),
173 "{body}"
174 );
175 }
176 }
177
178 #[tokio::test]
179 async fn refresh_success_parses_three_tokens() {
180 let mut server = mockito::Server::new_async().await;
181 server
182 .mock("POST", "/oauth/token")
183 .with_status(200)
184 .with_body(
185 r#"{"access_token":"new-at","refresh_token":"new-rt","id_token":"new-id","expires_in":3600}"#,
186 )
187 .create_async()
188 .await;
189 let client = reqwest::Client::new();
190 let r = refresh(&client, &format!("{}/oauth/token", server.url()), "old")
191 .await
192 .unwrap();
193 assert_eq!(r.access_token, "new-at");
194 assert_eq!(r.refresh_token.as_deref(), Some("new-rt"));
195 assert_eq!(r.id_token.as_deref(), Some("new-id"));
196 assert_eq!(r.expires_in, Some(3600));
197 }
198
199 #[tokio::test]
200 async fn malformed_success_does_not_echo_tokens_in_the_error() {
201 let mut server = mockito::Server::new_async().await;
202 server
203 .mock("POST", "/oauth/token")
204 .with_status(200)
205 .with_body(
206 r#"{"access_token":"sensitive-access-token","refresh_token":"sensitive-refresh-token","id_token":"sensitive-id-token","expires_in":"sensitive-schema-value"}"#,
207 )
208 .create_async()
209 .await;
210
211 let client = reqwest::Client::new();
212 let error = refresh(&client, &format!("{}/oauth/token", server.url()), "old")
213 .await
214 .unwrap_err()
215 .to_string();
216
217 assert!(error.contains("openai token response"));
218 assert!(!error.contains("sensitive-access-token"));
219 assert!(!error.contains("sensitive-refresh-token"));
220 assert!(!error.contains("sensitive-id-token"));
221 assert!(!error.contains("sensitive-schema-value"));
222 }
223
224 #[tokio::test]
225 async fn refresh_400_returns_http_with_description() {
226 let mut server = mockito::Server::new_async().await;
227 server
228 .mock("POST", "/oauth/token")
229 .with_status(400)
230 .with_body(r#"{"error":"invalid_grant","error_description":"Refresh expired"}"#)
231 .create_async()
232 .await;
233 let client = reqwest::Client::new();
234 let err = refresh(&client, &format!("{}/oauth/token", server.url()), "x")
235 .await
236 .unwrap_err();
237 match err {
238 AppError::Http { status, body } => {
239 assert_eq!(status, 400);
240 assert_eq!(body, "Refresh expired");
241 }
242 other => panic!("expected Http error, got {other:?}"),
243 }
244 }
245}