Skip to main content

pidge_client/auth/
device_code.rs

1//! RFC 8628 OAuth 2.0 device authorization grant flow.
2
3use std::time::Duration;
4
5use chrono::Utc;
6use serde::Deserialize;
7
8use crate::auth::tokens::TokenSet;
9use crate::error::ClientError;
10
11/// Response from the devicecode endpoint.
12#[derive(Debug, Clone, Deserialize)]
13pub struct DeviceCodeResponse {
14    pub device_code: String,
15    pub user_code: String,
16    pub verification_uri: String,
17    pub expires_in: u64,
18    pub interval: u64,
19    #[serde(default)]
20    pub message: Option<String>,
21}
22
23/// Token response. `refresh_token` is optional because refresh-token-rotation
24/// responses sometimes omit it; for device-code initial responses Microsoft always
25/// includes it (the `offline_access` scope is requested).
26#[derive(Debug, Deserialize)]
27pub(crate) struct TokenResponse {
28    pub access_token: String,
29    #[serde(default)]
30    pub refresh_token: Option<String>,
31    #[serde(default)]
32    pub expires_in: Option<u64>,
33    #[serde(default)]
34    pub id_token: Option<String>,
35}
36
37#[derive(Debug, Deserialize)]
38pub(crate) struct ErrorResponse {
39    pub error: String,
40    #[serde(default)]
41    pub error_description: Option<String>,
42}
43
44/// Result of a successful poll: tokens + the id_token (for tenant_id extraction).
45#[derive(Debug)]
46pub struct PollSuccess {
47    pub tokens: TokenSet,
48    pub id_token: Option<String>,
49}
50
51/// POST to `{base_url}/oauth2/v2.0/devicecode`.
52pub async fn start(
53    client: &reqwest::Client,
54    base_url: &str,
55    client_id: &str,
56    scope: &str,
57) -> Result<DeviceCodeResponse, ClientError> {
58    let url = format!("{base_url}/oauth2/v2.0/devicecode");
59    let resp = client
60        .post(&url)
61        .form(&[("client_id", client_id), ("scope", scope)])
62        .send()
63        .await?;
64
65    let status = resp.status();
66    if !status.is_success() {
67        let text = resp.text().await.unwrap_or_default();
68        return Err(ClientError::Graph {
69            status: status.as_u16(),
70            message: text,
71        });
72    }
73
74    Ok(resp.json::<DeviceCodeResponse>().await?)
75}
76
77/// Poll `{base_url}/oauth2/v2.0/token` per RFC 8628 ยง3.5.
78/// Returns when the user has approved consent, or with an error on denial / expiry.
79///
80/// The `sleep` parameter is a closure so tests can inject a fast no-op sleep
81/// instead of real `tokio::time::sleep`.
82pub async fn poll<F, Fut>(
83    client: &reqwest::Client,
84    base_url: &str,
85    client_id: &str,
86    device_code: &str,
87    initial_interval: u64,
88    expires_in: u64,
89    mut sleep: F,
90) -> Result<PollSuccess, ClientError>
91where
92    F: FnMut(Duration) -> Fut,
93    Fut: std::future::Future<Output = ()>,
94{
95    let url = format!("{base_url}/oauth2/v2.0/token");
96    let mut interval = initial_interval;
97    let deadline = Utc::now() + chrono::Duration::seconds(expires_in as i64);
98
99    loop {
100        if Utc::now() >= deadline {
101            return Err(ClientError::DeviceCodeTimeout);
102        }
103        sleep(Duration::from_secs(interval)).await;
104
105        let resp = client
106            .post(&url)
107            .form(&[
108                ("grant_type", "urn:ietf:params:oauth:grant-type:device_code"),
109                ("client_id", client_id),
110                ("device_code", device_code),
111            ])
112            .send()
113            .await?;
114
115        let status = resp.status();
116        let body = resp.bytes().await?;
117
118        if status.is_success() {
119            let tr: TokenResponse = serde_json::from_slice(&body)?;
120            let expires = tr.expires_in.unwrap_or(3600);
121            let refresh_token = tr.refresh_token.ok_or(ClientError::MissingAccessToken)?;
122            return Ok(PollSuccess {
123                tokens: TokenSet {
124                    access_token: tr.access_token,
125                    refresh_token,
126                    expires_at: Utc::now() + chrono::Duration::seconds(expires as i64 - 60),
127                },
128                id_token: tr.id_token,
129            });
130        }
131
132        let err: ErrorResponse = serde_json::from_slice(&body).map_err(|_| ClientError::Graph {
133            status: status.as_u16(),
134            message: String::from_utf8_lossy(&body).into_owned(),
135        })?;
136
137        match err.error.as_str() {
138            "authorization_pending" => continue,
139            "slow_down" => {
140                interval += 5;
141                continue;
142            }
143            "access_denied" => return Err(ClientError::DeviceCodeAccessDenied),
144            "expired_token" => return Err(ClientError::DeviceCodeTimeout),
145            other => {
146                return Err(ClientError::DeviceCodeOther {
147                    kind: other.to_string(),
148                    description: err.error_description,
149                });
150            }
151        }
152    }
153}
154
155/// Real-time sleep helper for production use.
156pub async fn real_sleep(d: Duration) {
157    tokio::time::sleep(d).await;
158}
159
160#[cfg(test)]
161mod tests {
162    use super::*;
163    use wiremock::matchers::{method, path};
164    use wiremock::{Mock, MockServer, ResponseTemplate};
165
166    /// Test sleep: no-op; we don't actually want to wait in tests.
167    async fn no_sleep(_: Duration) {}
168
169    #[tokio::test]
170    async fn start_returns_device_code_response() {
171        let server = MockServer::start().await;
172        Mock::given(method("POST"))
173            .and(path("/oauth2/v2.0/devicecode"))
174            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
175                "device_code": "DC123",
176                "user_code": "ABCD-1234",
177                "verification_uri": "https://microsoft.com/devicelogin",
178                "expires_in": 900,
179                "interval": 5,
180                "message": "Please sign in"
181            })))
182            .mount(&server)
183            .await;
184
185        let client = reqwest::Client::new();
186        let resp = start(&client, &server.uri(), "CID", "openid")
187            .await
188            .unwrap();
189        assert_eq!(resp.user_code, "ABCD-1234");
190        assert_eq!(resp.interval, 5);
191    }
192
193    #[tokio::test]
194    async fn poll_returns_tokens_on_success() {
195        let server = MockServer::start().await;
196        Mock::given(method("POST"))
197            .and(path("/oauth2/v2.0/token"))
198            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
199                "access_token": "AT123",
200                "refresh_token": "RT123",
201                "expires_in": 3600,
202                "id_token": "eyJh.eyJ0aWQiOiJ0aWQifQ.sig"
203            })))
204            .mount(&server)
205            .await;
206
207        let client = reqwest::Client::new();
208        let result = poll(&client, &server.uri(), "CID", "DC", 1, 60, no_sleep)
209            .await
210            .unwrap();
211        assert_eq!(result.tokens.access_token, "AT123");
212        assert_eq!(result.tokens.refresh_token, "RT123");
213        assert_eq!(
214            result.id_token.as_deref(),
215            Some("eyJh.eyJ0aWQiOiJ0aWQifQ.sig")
216        );
217    }
218
219    #[tokio::test]
220    async fn poll_retries_on_authorization_pending_then_succeeds() {
221        let server = MockServer::start().await;
222
223        Mock::given(method("POST"))
224            .and(path("/oauth2/v2.0/token"))
225            .respond_with(ResponseTemplate::new(400).set_body_json(serde_json::json!({
226                "error": "authorization_pending"
227            })))
228            .up_to_n_times(2)
229            .mount(&server)
230            .await;
231
232        Mock::given(method("POST"))
233            .and(path("/oauth2/v2.0/token"))
234            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
235                "access_token": "AT",
236                "refresh_token": "RT",
237                "expires_in": 3600
238            })))
239            .mount(&server)
240            .await;
241
242        let client = reqwest::Client::new();
243        let result = poll(&client, &server.uri(), "CID", "DC", 0, 60, no_sleep)
244            .await
245            .unwrap();
246        assert_eq!(result.tokens.access_token, "AT");
247    }
248
249    #[tokio::test]
250    async fn poll_returns_access_denied_on_user_cancel() {
251        let server = MockServer::start().await;
252        Mock::given(method("POST"))
253            .and(path("/oauth2/v2.0/token"))
254            .respond_with(ResponseTemplate::new(400).set_body_json(serde_json::json!({
255                "error": "access_denied"
256            })))
257            .mount(&server)
258            .await;
259
260        let client = reqwest::Client::new();
261        let result = poll(&client, &server.uri(), "CID", "DC", 0, 60, no_sleep).await;
262        assert!(matches!(result, Err(ClientError::DeviceCodeAccessDenied)));
263    }
264
265    #[tokio::test]
266    async fn poll_returns_timeout_on_expired_token() {
267        let server = MockServer::start().await;
268        Mock::given(method("POST"))
269            .and(path("/oauth2/v2.0/token"))
270            .respond_with(ResponseTemplate::new(400).set_body_json(serde_json::json!({
271                "error": "expired_token"
272            })))
273            .mount(&server)
274            .await;
275
276        let client = reqwest::Client::new();
277        let result = poll(&client, &server.uri(), "CID", "DC", 0, 60, no_sleep).await;
278        assert!(matches!(result, Err(ClientError::DeviceCodeTimeout)));
279    }
280
281    #[tokio::test]
282    async fn poll_returns_other_on_unknown_error_code() {
283        let server = MockServer::start().await;
284        Mock::given(method("POST"))
285            .and(path("/oauth2/v2.0/token"))
286            .respond_with(ResponseTemplate::new(400).set_body_json(serde_json::json!({
287                "error": "consent_required",
288                "error_description": "AADSTS65001: consent needed"
289            })))
290            .mount(&server)
291            .await;
292
293        let client = reqwest::Client::new();
294        let result = poll(&client, &server.uri(), "CID", "DC", 0, 60, no_sleep).await;
295        match result {
296            Err(ClientError::DeviceCodeOther { kind, description }) => {
297                assert_eq!(kind, "consent_required");
298                assert_eq!(description.as_deref(), Some("AADSTS65001: consent needed"));
299            }
300            other => panic!("expected DeviceCodeOther, got {other:?}"),
301        }
302    }
303
304    #[tokio::test]
305    async fn poll_slow_down_increases_interval_then_succeeds() {
306        let server = MockServer::start().await;
307
308        Mock::given(method("POST"))
309            .and(path("/oauth2/v2.0/token"))
310            .respond_with(ResponseTemplate::new(400).set_body_json(serde_json::json!({
311                "error": "slow_down"
312            })))
313            .up_to_n_times(1)
314            .mount(&server)
315            .await;
316
317        Mock::given(method("POST"))
318            .and(path("/oauth2/v2.0/token"))
319            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
320                "access_token": "AT",
321                "refresh_token": "RT",
322                "expires_in": 3600
323            })))
324            .mount(&server)
325            .await;
326
327        let client = reqwest::Client::new();
328        let result = poll(&client, &server.uri(), "CID", "DC", 0, 60, no_sleep)
329            .await
330            .unwrap();
331        assert_eq!(result.tokens.access_token, "AT");
332    }
333}