Skip to main content

redisctl_core/auth/
device_flow.rs

1//! OIDC Device Authorization Grant (RFC 8628) client for the Redis Cloud Okta tenant.
2//!
3//! Two requests against the issuer, handled by the [`oauth2`] crate:
4//! - `POST {issuer}/v1/device/authorize` — [`start`](DeviceFlowClient::start); returns the user
5//!   code + verification URI plus the (secret) device code needed to resume.
6//! - `POST {issuer}/v1/token` (device-code grant) — [`poll`](DeviceFlowClient::poll); the crate
7//!   owns the polling loop (honoring the server's `interval`/`slow_down`) until the request is
8//!   approved, denied, or times out.
9//!
10//! `login --device` is deliberately split: the CLI calls [`start`](DeviceFlowClient::start),
11//! prints/persists the returned [`DeviceAuthorization`] (which is `serde`-serializable), and
12//! returns. A later `status --wait` — possibly a *different* process — rebuilds the client from
13//! the issuer + client id, deserializes the [`DeviceAuthorization`], and calls
14//! [`poll`](DeviceFlowClient::poll). Refreshing a token is flow-agnostic and lives on
15//! [`CloudAuthenticator::refresh`](super::authenticator::CloudAuthenticator::refresh).
16
17use 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/// Device-authorization-grant client bound to one Okta issuer + public client id.
29///
30/// Cheap to construct and holds no network state, so `status --wait` can rebuild it from the
31/// persisted issuer + client id to resume a login started by an earlier `login --device`.
32#[derive(Clone)]
33pub struct DeviceFlowClient {
34    issuer: Url,
35    client_id: String,
36}
37
38/// The device-authorization response — what `auth login --device` surfaces to the developer and
39/// persists so a later poll (in any process) can resume.
40///
41/// Wraps the RFC 8628 response from the [`oauth2`] crate. It is `serde`-serializable so the CLI
42/// can persist it between the non-blocking `login --device` and a later `status --wait`. The
43/// `device_code` it carries is a secret; `Debug` redacts it (as do the other code fields) via
44/// the underlying `oauth2` secret types.
45#[derive(Clone, Debug, Serialize, Deserialize)]
46pub struct DeviceAuthorization {
47    inner: StandardDeviceAuthorizationResponse,
48}
49
50impl DeviceAuthorization {
51    /// The end-user verification code (shown to the user to type at the verification URI).
52    pub fn user_code(&self) -> &str {
53        self.inner.user_code().secret().as_str()
54    }
55
56    /// The page the user opens and then *types* the `user_code` into.
57    pub fn verification_uri(&self) -> &str {
58        self.inner.verification_uri().as_str()
59    }
60
61    /// The same page with the `user_code` pre-embedded (optional per RFC 8628), so the user can
62    /// just open it and confirm — no manual code entry. Prefer this when present.
63    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    /// Lifetime of the device/user code, in seconds.
70    pub fn expires_in(&self) -> u64 {
71        self.inner.expires_in().as_secs()
72    }
73
74    /// Minimum poll interval requested by the server, in seconds.
75    pub fn interval(&self) -> u64 {
76        self.inner.interval().as_secs()
77    }
78
79    /// Borrow the underlying RFC 8628 response (escape hatch for callers that need the raw type).
80    pub fn as_standard(&self) -> &StandardDeviceAuthorizationResponse {
81        &self.inner
82    }
83}
84
85impl DeviceFlowClient {
86    /// Build a client for the given issuer (e.g.
87    /// `https://<your-okta-issuer>/oauth2/default`) and public client id.
88    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    /// Start device authorization: `POST /v1/device/authorize`.
96    ///
97    /// Returns the codes to display *and* the device code needed to resume polling — persist the
98    /// returned [`DeviceAuthorization`] and hand it to [`poll`](Self::poll) later.
99    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    /// Poll the token endpoint until the user approves, the code is denied/expires, or `timeout`
114    /// elapses. The [`oauth2`] crate runs the loop internally (respecting the server's poll
115    /// interval and `slow_down`), so this is a single blocking call.
116    ///
117    /// `timeout` bounds the whole wait; `None` falls back to the device code's own lifetime. A
118    /// timeout surfaces as [`AuthError::Expired`]. Callable from a freshly built client after
119    /// deserializing `authz`.
120    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        // Debug must not leak the secret device code / user code.
191        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        // RFC 8628 default poll interval when the server omits one.
209        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    /// The device authorization must survive a serialize/deserialize round-trip and still be
287    /// usable to poll — this is the `login --device` (persist) then `status --wait` (resume,
288    /// possibly in another process) contract.
289    #[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        // A freshly built client (as a separate process would build) can complete the poll.
312        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}