Skip to main content

claude_codex/providers/codex/auth/
jwt.rs

1use serde::Deserialize;
2
3#[derive(Debug, Deserialize)]
4struct IdTokenClaims {
5    #[serde(default)]
6    chatgpt_account_id: Option<String>,
7    #[serde(default)]
8    organizations: Option<Vec<OrgClaim>>,
9    #[allow(dead_code)]
10    #[serde(default)]
11    email: Option<String>,
12    #[serde(default)]
13    #[serde(rename = "https://api.openai.com/auth")]
14    openai_auth: Option<OpenAiAuthClaim>,
15    #[serde(default)]
16    #[serde(rename = "https://api.openai.com/auth.chatgpt_account_id")]
17    openai_chatgpt_account_id: Option<String>,
18}
19
20#[derive(Debug, Deserialize)]
21struct OrgClaim {
22    id: String,
23}
24
25#[derive(Debug, Deserialize)]
26struct OpenAiAuthClaim {
27    #[serde(default)]
28    chatgpt_account_id: Option<String>,
29}
30
31#[derive(Debug, Deserialize)]
32pub struct TokenResponse {
33    pub id_token: Option<String>,
34    pub access_token: String,
35    pub refresh_token: String,
36    pub expires_in: Option<u64>,
37}
38
39fn decode_jwt_payload(token: &str) -> Option<Vec<u8>> {
40    let parts: Vec<&str> = token.split('.').collect();
41    if parts.len() != 3 {
42        return None;
43    }
44    let payload_b64 = parts[1].replace('-', "+").replace('_', "/");
45    let padded = match payload_b64.len() % 4 {
46        2 => format!("{payload_b64}=="),
47        3 => format!("{payload_b64}="),
48        _ => payload_b64,
49    };
50    use base64::Engine;
51    base64::engine::general_purpose::STANDARD
52        .decode(&padded)
53        .ok()
54}
55
56fn parse_jwt_claims(token: &str) -> Option<IdTokenClaims> {
57    serde_json::from_slice(&decode_jwt_payload(token)?).ok()
58}
59
60/// Decode a JWT's `exp` claim (seconds since the Unix epoch) into epoch milliseconds.
61pub fn token_exp_ms(token: &str) -> Option<u64> {
62    #[derive(Deserialize)]
63    struct ExpClaim {
64        exp: Option<u64>,
65    }
66    let claim: ExpClaim = serde_json::from_slice(&decode_jwt_payload(token)?).ok()?;
67    claim.exp.map(|secs| secs.saturating_mul(1000))
68}
69
70fn extract_account_id_from_claims(claims: &IdTokenClaims) -> Option<String> {
71    claims
72        .chatgpt_account_id
73        .clone()
74        .or_else(|| claims.openai_auth.as_ref()?.chatgpt_account_id.clone())
75        .or_else(|| claims.openai_chatgpt_account_id.clone())
76        .or_else(|| claims.organizations.as_ref()?.first()?.id.clone().into())
77}
78
79pub fn validate_token_response(tokens: &TokenResponse) -> anyhow::Result<()> {
80    if tokens.access_token.trim().is_empty() {
81        anyhow::bail!("token response missing access token");
82    }
83    if tokens.refresh_token.trim().is_empty() {
84        anyhow::bail!("token response missing refresh token");
85    }
86    if matches!(tokens.expires_in, Some(0)) {
87        anyhow::bail!("token response has invalid expiration");
88    }
89    Ok(())
90}
91
92pub fn extract_account_id(tokens: &TokenResponse) -> Option<String> {
93    if let Some(ref id_token) = tokens.id_token
94        && let Some(claims) = parse_jwt_claims(id_token)
95        && let Some(account_id) = extract_account_id_from_claims(&claims)
96    {
97        return Some(account_id);
98    }
99    let claims = parse_jwt_claims(&tokens.access_token)?;
100    extract_account_id_from_claims(&claims)
101}
102
103#[cfg(test)]
104mod tests {
105    use super::*;
106
107    #[test]
108    fn extract_account_id_from_access_token() {
109        let token = TokenResponse {
110            id_token: None,
111            access_token: "eyJhbGciOiJIUzI1NiJ9.eyJjaGF0Z3B0X2FjY291bnRfaWQiOiJhY2N0XzEyMyJ9.sig"
112                .into(),
113            refresh_token: "r".into(),
114            expires_in: Some(3600),
115        };
116        assert_eq!(extract_account_id(&token), Some("acct_123".into()));
117    }
118
119    #[test]
120    fn extract_account_id_from_id_token_takes_precedence() {
121        let token = TokenResponse {
122            id_token: Some(
123                "eyJhbGciOiJIUzI1NiJ9.eyJjaGF0Z3B0X2FjY291bnRfaWQiOiJpZF9hY2N0In0.sig".into(),
124            ),
125            access_token: "eyJhbGciOiJIUzI1NiJ9.eyJjaGF0Z3B0X2FjY291bnRfaWQiOiJhY2NfYWNjIn0.sig"
126                .into(),
127            refresh_token: "r".into(),
128            expires_in: Some(3600),
129        };
130        assert_eq!(extract_account_id(&token), Some("id_acct".into()));
131    }
132
133    #[test]
134    fn extract_account_id_returns_none_for_invalid_token() {
135        let token = TokenResponse {
136            id_token: None,
137            access_token: "invalid".into(),
138            refresh_token: "r".into(),
139            expires_in: None,
140        };
141        assert_eq!(extract_account_id(&token), None);
142    }
143
144    #[test]
145    fn validate_token_response_rejects_empty_access_token() {
146        let token = TokenResponse {
147            access_token: "".into(),
148            refresh_token: "r".into(),
149            expires_in: Some(3600),
150            id_token: None,
151        };
152        assert!(validate_token_response(&token).is_err());
153        assert!(
154            validate_token_response(&token)
155                .unwrap_err()
156                .to_string()
157                .contains("missing access token")
158        );
159    }
160
161    #[test]
162    fn validate_token_response_rejects_empty_refresh_token() {
163        let token = TokenResponse {
164            access_token: "a".into(),
165            refresh_token: "".into(),
166            expires_in: Some(3600),
167            id_token: None,
168        };
169        assert!(validate_token_response(&token).is_err());
170        assert!(
171            validate_token_response(&token)
172                .unwrap_err()
173                .to_string()
174                .contains("missing refresh token")
175        );
176    }
177
178    #[test]
179    fn validate_token_response_rejects_zero_expires_in() {
180        let token = TokenResponse {
181            access_token: "a".into(),
182            refresh_token: "r".into(),
183            expires_in: Some(0),
184            id_token: None,
185        };
186        assert!(validate_token_response(&token).is_err());
187        assert!(
188            validate_token_response(&token)
189                .unwrap_err()
190                .to_string()
191                .contains("invalid expiration")
192        );
193    }
194
195    #[test]
196    fn validate_token_response_accepts_valid() {
197        let token = TokenResponse {
198            access_token: "a".into(),
199            refresh_token: "r".into(),
200            expires_in: Some(3600),
201            id_token: None,
202        };
203        assert!(validate_token_response(&token).is_ok());
204    }
205
206    #[test]
207    fn validate_token_response_accepts_no_expires_in() {
208        let token = TokenResponse {
209            access_token: "a".into(),
210            refresh_token: "r".into(),
211            expires_in: None,
212            id_token: None,
213        };
214        assert!(validate_token_response(&token).is_ok());
215    }
216}