Skip to main content

pidge_client/auth/
refresh.rs

1//! OAuth refresh-token grant.
2
3use chrono::Utc;
4
5use crate::auth::device_code::{ErrorResponse, TokenResponse};
6use crate::auth::tokens::TokenSet;
7use crate::error::ClientError;
8
9/// Refresh the access token using the stored refresh token.
10///
11/// Returns the new `TokenSet`. If Microsoft rotates the refresh token, it's
12/// included in the response and used. If not, the current refresh token is preserved.
13///
14/// Errors:
15/// - `ClientError::SessionExpired` if Microsoft returns `invalid_grant`
16/// - `ClientError::Graph` for other HTTP errors
17pub async fn refresh(
18    client: &reqwest::Client,
19    base_url: &str,
20    client_id: &str,
21    current: &TokenSet,
22    scope: &str,
23    email: &str,
24) -> Result<TokenSet, ClientError> {
25    let url = format!("{base_url}/oauth2/v2.0/token");
26    let resp = client
27        .post(&url)
28        .form(&[
29            ("grant_type", "refresh_token"),
30            ("client_id", client_id),
31            ("refresh_token", &current.refresh_token),
32            ("scope", scope),
33        ])
34        .send()
35        .await?;
36
37    let status = resp.status();
38    let body = resp.bytes().await?;
39
40    if !status.is_success() {
41        let err: ErrorResponse = serde_json::from_slice(&body).map_err(|_| ClientError::Graph {
42            status: status.as_u16(),
43            message: String::from_utf8_lossy(&body).into_owned(),
44        })?;
45        if err.error == "invalid_grant" {
46            return Err(ClientError::SessionExpired {
47                email: email.to_string(),
48            });
49        }
50        return Err(ClientError::Graph {
51            status: status.as_u16(),
52            message: err.error_description.unwrap_or(err.error),
53        });
54    }
55
56    let tr: TokenResponse = serde_json::from_slice(&body)?;
57    let expires = tr.expires_in.unwrap_or(3600);
58    let new_refresh = tr
59        .refresh_token
60        .unwrap_or_else(|| current.refresh_token.clone());
61
62    Ok(TokenSet {
63        access_token: tr.access_token,
64        refresh_token: new_refresh,
65        expires_at: Utc::now() + chrono::Duration::seconds(expires as i64 - 60),
66    })
67}
68
69#[cfg(test)]
70mod tests {
71    use super::*;
72    use chrono::Duration;
73    use wiremock::matchers::{method, path};
74    use wiremock::{Mock, MockServer, ResponseTemplate};
75
76    fn old_tokens() -> TokenSet {
77        TokenSet {
78            access_token: "OLD_AT".into(),
79            refresh_token: "OLD_RT".into(),
80            expires_at: Utc::now() - Duration::seconds(60),
81        }
82    }
83
84    #[tokio::test]
85    async fn refresh_returns_new_tokens_on_success() {
86        let server = MockServer::start().await;
87        Mock::given(method("POST"))
88            .and(path("/oauth2/v2.0/token"))
89            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
90                "access_token": "NEW_AT",
91                "refresh_token": "NEW_RT",
92                "expires_in": 3600
93            })))
94            .mount(&server)
95            .await;
96
97        let client = reqwest::Client::new();
98        let new = refresh(
99            &client,
100            &server.uri(),
101            "CID",
102            &old_tokens(),
103            "scope",
104            "u@e.com",
105        )
106        .await
107        .unwrap();
108        assert_eq!(new.access_token, "NEW_AT");
109        assert_eq!(new.refresh_token, "NEW_RT");
110    }
111
112    #[tokio::test]
113    async fn refresh_preserves_refresh_token_when_response_omits_it() {
114        let server = MockServer::start().await;
115        Mock::given(method("POST"))
116            .and(path("/oauth2/v2.0/token"))
117            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
118                "access_token": "NEW_AT",
119                "expires_in": 3600
120            })))
121            .mount(&server)
122            .await;
123
124        let client = reqwest::Client::new();
125        let new = refresh(
126            &client,
127            &server.uri(),
128            "CID",
129            &old_tokens(),
130            "scope",
131            "u@e.com",
132        )
133        .await
134        .unwrap();
135        assert_eq!(new.access_token, "NEW_AT");
136        assert_eq!(new.refresh_token, "OLD_RT");
137    }
138
139    #[tokio::test]
140    async fn refresh_returns_session_expired_on_invalid_grant() {
141        let server = MockServer::start().await;
142        Mock::given(method("POST"))
143            .and(path("/oauth2/v2.0/token"))
144            .respond_with(ResponseTemplate::new(400).set_body_json(serde_json::json!({
145                "error": "invalid_grant",
146                "error_description": "AADSTS50173: refresh token expired"
147            })))
148            .mount(&server)
149            .await;
150
151        let client = reqwest::Client::new();
152        let err = refresh(
153            &client,
154            &server.uri(),
155            "CID",
156            &old_tokens(),
157            "scope",
158            "u@e.com",
159        )
160        .await
161        .unwrap_err();
162        match err {
163            ClientError::SessionExpired { email } => assert_eq!(email, "u@e.com"),
164            other => panic!("expected SessionExpired, got {other:?}"),
165        }
166    }
167}