1use std::time::Duration;
18
19use oauth2::{Scope, StandardDeviceAuthorizationResponse};
20use serde::{Deserialize, Serialize};
21use url::Url;
22
23use super::oidc::{
24 map_basic_token_error, map_device_token_error, oauth_http_client, okta_client, to_token_set,
25};
26use super::{AuthError, TokenSet};
27
28#[derive(Clone)]
33pub struct DeviceFlowClient {
34 issuer: Url,
35 client_id: String,
36}
37
38#[derive(Clone, Debug, Serialize, Deserialize)]
46pub struct DeviceAuthorization {
47 inner: StandardDeviceAuthorizationResponse,
48}
49
50impl DeviceAuthorization {
51 pub fn user_code(&self) -> &str {
53 self.inner.user_code().secret().as_str()
54 }
55
56 pub fn verification_uri(&self) -> &str {
58 self.inner.verification_uri().as_str()
59 }
60
61 pub fn verification_uri_complete(&self) -> Option<&str> {
64 self.inner
65 .verification_uri_complete()
66 .map(|v| v.secret().as_str())
67 }
68
69 pub fn expires_in(&self) -> u64 {
71 self.inner.expires_in().as_secs()
72 }
73
74 pub fn interval(&self) -> u64 {
76 self.inner.interval().as_secs()
77 }
78
79 pub fn as_standard(&self) -> &StandardDeviceAuthorizationResponse {
81 &self.inner
82 }
83}
84
85impl DeviceFlowClient {
86 pub fn new(issuer: Url, client_id: impl Into<String>) -> Self {
89 Self {
90 issuer,
91 client_id: client_id.into(),
92 }
93 }
94
95 pub async fn start(&self, scopes: &[&str]) -> Result<DeviceAuthorization, AuthError> {
100 let client = okta_client(&self.issuer, &self.client_id)?;
101 let http = oauth_http_client()?;
102 let mut request = client.exchange_device_code();
103 for scope in scopes {
104 request = request.add_scope(Scope::new((*scope).to_string()));
105 }
106 let inner: StandardDeviceAuthorizationResponse = request
107 .request_async(&http)
108 .await
109 .map_err(map_basic_token_error)?;
110 Ok(DeviceAuthorization { inner })
111 }
112
113 pub async fn poll(
121 &self,
122 authz: &DeviceAuthorization,
123 timeout: Option<Duration>,
124 ) -> Result<TokenSet, AuthError> {
125 let client = okta_client(&self.issuer, &self.client_id)?;
126 let http = oauth_http_client()?;
127 let resp = client
128 .exchange_device_access_token(&authz.inner)
129 .request_async(&http, tokio::time::sleep, timeout)
130 .await
131 .map_err(map_device_token_error)?;
132 Ok(to_token_set(&resp))
133 }
134}
135
136#[cfg(test)]
137mod tests {
138 use super::*;
139 use wiremock::matchers::{method, path};
140 use wiremock::{Mock, MockServer, ResponseTemplate};
141
142 fn client(server: &MockServer) -> DeviceFlowClient {
143 DeviceFlowClient::new(Url::parse(&server.uri()).unwrap(), "test-client")
144 }
145
146 async fn mount_device_authorize(server: &MockServer, body: serde_json::Value) {
147 Mock::given(method("POST"))
148 .and(path("/v1/device/authorize"))
149 .respond_with(ResponseTemplate::new(200).set_body_json(body))
150 .mount(server)
151 .await;
152 }
153
154 async fn mount_token(server: &MockServer, status: u16, body: serde_json::Value) {
155 Mock::given(method("POST"))
156 .and(path("/v1/token"))
157 .respond_with(ResponseTemplate::new(status).set_body_json(body))
158 .mount(server)
159 .await;
160 }
161
162 #[tokio::test]
163 async fn start_parses_device_authorization() {
164 let server = MockServer::start().await;
165 mount_device_authorize(
166 &server,
167 serde_json::json!({
168 "device_code": "DC",
169 "user_code": "WDJB-MJHT",
170 "verification_uri": "https://x/activate",
171 "verification_uri_complete": "https://x/activate?user_code=WDJB-MJHT",
172 "expires_in": 600,
173 "interval": 5
174 }),
175 )
176 .await;
177
178 let d = client(&server)
179 .start(&["openid", "offline_access"])
180 .await
181 .unwrap();
182 assert_eq!(d.user_code(), "WDJB-MJHT");
183 assert_eq!(d.verification_uri(), "https://x/activate");
184 assert_eq!(
185 d.verification_uri_complete(),
186 Some("https://x/activate?user_code=WDJB-MJHT")
187 );
188 assert_eq!(d.expires_in(), 600);
189 assert_eq!(d.interval(), 5);
190 assert!(!format!("{d:?}").contains("WDJB-MJHT"));
192 }
193
194 #[tokio::test]
195 async fn start_defaults_interval_when_absent() {
196 let server = MockServer::start().await;
197 mount_device_authorize(
198 &server,
199 serde_json::json!({
200 "device_code": "DC",
201 "user_code": "U",
202 "verification_uri": "https://x",
203 "expires_in": 600
204 }),
205 )
206 .await;
207
208 let d = client(&server).start(&["openid"]).await.unwrap();
210 assert_eq!(d.interval(), 5);
211 }
212
213 #[tokio::test]
214 async fn poll_ready_returns_tokens() {
215 let server = MockServer::start().await;
216 mount_device_authorize(
217 &server,
218 serde_json::json!({
219 "device_code": "DC", "user_code": "U",
220 "verification_uri": "https://x", "expires_in": 600, "interval": 5
221 }),
222 )
223 .await;
224 mount_token(
225 &server,
226 200,
227 serde_json::json!({
228 "access_token": "AT",
229 "token_type": "Bearer",
230 "refresh_token": "RT",
231 "expires_in": 3600
232 }),
233 )
234 .await;
235
236 let c = client(&server);
237 let authz = c.start(&["openid"]).await.unwrap();
238 let t = c.poll(&authz, Some(Duration::from_secs(5))).await.unwrap();
239 assert_eq!(t.access_token, "AT");
240 assert_eq!(t.refresh_token.as_deref(), Some("RT"));
241 assert_eq!(t.expires_in, 3600);
242 }
243
244 #[tokio::test]
245 async fn poll_expired_maps_to_error() {
246 let server = MockServer::start().await;
247 mount_device_authorize(
248 &server,
249 serde_json::json!({
250 "device_code": "DC", "user_code": "U",
251 "verification_uri": "https://x", "expires_in": 600, "interval": 5
252 }),
253 )
254 .await;
255 mount_token(&server, 400, serde_json::json!({"error": "expired_token"})).await;
256
257 let c = client(&server);
258 let authz = c.start(&["openid"]).await.unwrap();
259 assert!(matches!(
260 c.poll(&authz, Some(Duration::from_secs(5))).await,
261 Err(AuthError::Expired)
262 ));
263 }
264
265 #[tokio::test]
266 async fn poll_denied_maps_to_error() {
267 let server = MockServer::start().await;
268 mount_device_authorize(
269 &server,
270 serde_json::json!({
271 "device_code": "DC", "user_code": "U",
272 "verification_uri": "https://x", "expires_in": 600, "interval": 5
273 }),
274 )
275 .await;
276 mount_token(&server, 400, serde_json::json!({"error": "access_denied"})).await;
277
278 let c = client(&server);
279 let authz = c.start(&["openid"]).await.unwrap();
280 assert!(matches!(
281 c.poll(&authz, Some(Duration::from_secs(5))).await,
282 Err(AuthError::Denied)
283 ));
284 }
285
286 #[tokio::test]
290 async fn device_authorization_round_trips_and_resumes() {
291 let server = MockServer::start().await;
292 mount_device_authorize(
293 &server,
294 serde_json::json!({
295 "device_code": "DC", "user_code": "U",
296 "verification_uri": "https://x", "expires_in": 600, "interval": 5
297 }),
298 )
299 .await;
300 mount_token(
301 &server,
302 200,
303 serde_json::json!({"access_token": "AT", "token_type": "Bearer", "expires_in": 3600}),
304 )
305 .await;
306
307 let started = client(&server).start(&["openid"]).await.unwrap();
308 let json = serde_json::to_string(&started).unwrap();
309 let resumed: DeviceAuthorization = serde_json::from_str(&json).unwrap();
310
311 let fresh = DeviceFlowClient::new(Url::parse(&server.uri()).unwrap(), "test-client");
313 let t = fresh
314 .poll(&resumed, Some(Duration::from_secs(5)))
315 .await
316 .unwrap();
317 assert_eq!(t.access_token, "AT");
318 }
319}