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.decode(&padded).ok()
52}
53
54fn parse_jwt_claims(token: &str) -> Option<IdTokenClaims> {
55    serde_json::from_slice(&decode_jwt_payload(token)?).ok()
56}
57
58/// Decode a JWT's `exp` claim (seconds since the Unix epoch) into epoch milliseconds.
59pub fn token_exp_ms(token: &str) -> Option<u64> {
60    #[derive(Deserialize)]
61    struct ExpClaim {
62        exp: Option<u64>,
63    }
64    let claim: ExpClaim = serde_json::from_slice(&decode_jwt_payload(token)?).ok()?;
65    claim.exp.map(|secs| secs.saturating_mul(1000))
66}
67
68fn extract_account_id_from_claims(claims: &IdTokenClaims) -> Option<String> {
69    claims
70        .chatgpt_account_id
71        .clone()
72        .or_else(|| claims.openai_auth.as_ref()?.chatgpt_account_id.clone())
73        .or_else(|| claims.openai_chatgpt_account_id.clone())
74        .or_else(|| claims.organizations.as_ref()?.first()?.id.clone().into())
75}
76
77pub fn validate_token_response(tokens: &TokenResponse) -> anyhow::Result<()> {
78    if tokens.access_token.trim().is_empty() {
79        anyhow::bail!("token response missing access token");
80    }
81    if tokens.refresh_token.trim().is_empty() {
82        anyhow::bail!("token response missing refresh token");
83    }
84    if matches!(tokens.expires_in, Some(0)) {
85        anyhow::bail!("token response has invalid expiration");
86    }
87    Ok(())
88}
89
90pub fn extract_account_id(tokens: &TokenResponse) -> Option<String> {
91    if let Some(ref id_token) = tokens.id_token
92        && let Some(claims) = parse_jwt_claims(id_token)
93        && let Some(account_id) = extract_account_id_from_claims(&claims)
94    {
95        return Some(account_id);
96    }
97    let claims = parse_jwt_claims(&tokens.access_token)?;
98    extract_account_id_from_claims(&claims)
99}
100
101#[cfg(test)]
102mod tests {
103    use super::*;
104
105    #[test]
106    fn extract_account_id_from_access_token() {
107        let token = TokenResponse {
108            id_token: None,
109            access_token: "eyJhbGciOiJIUzI1NiJ9.eyJjaGF0Z3B0X2FjY291bnRfaWQiOiJhY2N0XzEyMyJ9.sig"
110                .into(),
111            refresh_token: "r".into(),
112            expires_in: Some(3600),
113        };
114        assert_eq!(extract_account_id(&token), Some("acct_123".into()));
115    }
116
117    #[test]
118    fn extract_account_id_from_id_token_takes_precedence() {
119        let token = TokenResponse {
120            id_token: Some(
121                "eyJhbGciOiJIUzI1NiJ9.eyJjaGF0Z3B0X2FjY291bnRfaWQiOiJpZF9hY2N0In0.sig".into(),
122            ),
123            access_token: "eyJhbGciOiJIUzI1NiJ9.eyJjaGF0Z3B0X2FjY291bnRfaWQiOiJhY2NfYWNjIn0.sig"
124                .into(),
125            refresh_token: "r".into(),
126            expires_in: Some(3600),
127        };
128        assert_eq!(extract_account_id(&token), Some("id_acct".into()));
129    }
130
131    #[test]
132    fn extract_account_id_returns_none_for_invalid_token() {
133        let token = TokenResponse {
134            id_token: None,
135            access_token: "invalid".into(),
136            refresh_token: "r".into(),
137            expires_in: None,
138        };
139        assert_eq!(extract_account_id(&token), None);
140    }
141
142    #[test]
143    fn validate_token_response_rejects_empty_access_token() {
144        let token = TokenResponse {
145            access_token: "".into(),
146            refresh_token: "r".into(),
147            expires_in: Some(3600),
148            id_token: None,
149        };
150        assert!(validate_token_response(&token).is_err());
151        assert!(
152            validate_token_response(&token)
153                .unwrap_err()
154                .to_string()
155                .contains("missing access token")
156        );
157    }
158
159    #[test]
160    fn validate_token_response_rejects_empty_refresh_token() {
161        let token = TokenResponse {
162            access_token: "a".into(),
163            refresh_token: "".into(),
164            expires_in: Some(3600),
165            id_token: None,
166        };
167        assert!(validate_token_response(&token).is_err());
168        assert!(
169            validate_token_response(&token)
170                .unwrap_err()
171                .to_string()
172                .contains("missing refresh token")
173        );
174    }
175
176    #[test]
177    fn validate_token_response_rejects_zero_expires_in() {
178        let token = TokenResponse {
179            access_token: "a".into(),
180            refresh_token: "r".into(),
181            expires_in: Some(0),
182            id_token: None,
183        };
184        assert!(validate_token_response(&token).is_err());
185        assert!(
186            validate_token_response(&token)
187                .unwrap_err()
188                .to_string()
189                .contains("invalid expiration")
190        );
191    }
192
193    #[test]
194    fn validate_token_response_accepts_valid() {
195        let token = TokenResponse {
196            access_token: "a".into(),
197            refresh_token: "r".into(),
198            expires_in: Some(3600),
199            id_token: None,
200        };
201        assert!(validate_token_response(&token).is_ok());
202    }
203
204    #[test]
205    fn validate_token_response_accepts_no_expires_in() {
206        let token = TokenResponse {
207            access_token: "a".into(),
208            refresh_token: "r".into(),
209            expires_in: None,
210            id_token: None,
211        };
212        assert!(validate_token_response(&token).is_ok());
213    }
214}