1use std::time::Duration;
4
5use chrono::Utc;
6use serde::Deserialize;
7
8use crate::auth::tokens::TokenSet;
9use crate::error::ClientError;
10
11#[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#[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#[derive(Debug)]
46pub struct PollSuccess {
47 pub tokens: TokenSet,
48 pub id_token: Option<String>,
49}
50
51pub 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
77pub 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
155pub 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 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}