use std::time::Duration;
use oauth2::{Scope, StandardDeviceAuthorizationResponse};
use serde::{Deserialize, Serialize};
use url::Url;
use super::oidc::{
map_basic_token_error, map_device_token_error, oauth_http_client, okta_client, to_token_set,
};
use super::{AuthError, TokenSet};
#[derive(Clone)]
pub struct DeviceFlowClient {
issuer: Url,
client_id: String,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct DeviceAuthorization {
inner: StandardDeviceAuthorizationResponse,
}
impl DeviceAuthorization {
pub fn user_code(&self) -> &str {
self.inner.user_code().secret().as_str()
}
pub fn verification_uri(&self) -> &str {
self.inner.verification_uri().as_str()
}
pub fn verification_uri_complete(&self) -> Option<&str> {
self.inner
.verification_uri_complete()
.map(|v| v.secret().as_str())
}
pub fn expires_in(&self) -> u64 {
self.inner.expires_in().as_secs()
}
pub fn interval(&self) -> u64 {
self.inner.interval().as_secs()
}
pub fn as_standard(&self) -> &StandardDeviceAuthorizationResponse {
&self.inner
}
}
impl DeviceFlowClient {
pub fn new(issuer: Url, client_id: impl Into<String>) -> Self {
Self {
issuer,
client_id: client_id.into(),
}
}
pub async fn start(&self, scopes: &[&str]) -> Result<DeviceAuthorization, AuthError> {
let client = okta_client(&self.issuer, &self.client_id)?;
let http = oauth_http_client()?;
let mut request = client.exchange_device_code();
for scope in scopes {
request = request.add_scope(Scope::new((*scope).to_string()));
}
let inner: StandardDeviceAuthorizationResponse = request
.request_async(&http)
.await
.map_err(map_basic_token_error)?;
Ok(DeviceAuthorization { inner })
}
pub async fn poll(
&self,
authz: &DeviceAuthorization,
timeout: Option<Duration>,
) -> Result<TokenSet, AuthError> {
let client = okta_client(&self.issuer, &self.client_id)?;
let http = oauth_http_client()?;
let resp = client
.exchange_device_access_token(&authz.inner)
.request_async(&http, tokio::time::sleep, timeout)
.await
.map_err(map_device_token_error)?;
Ok(to_token_set(&resp))
}
}
#[cfg(test)]
mod tests {
use super::*;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn client(server: &MockServer) -> DeviceFlowClient {
DeviceFlowClient::new(Url::parse(&server.uri()).unwrap(), "test-client")
}
async fn mount_device_authorize(server: &MockServer, body: serde_json::Value) {
Mock::given(method("POST"))
.and(path("/v1/device/authorize"))
.respond_with(ResponseTemplate::new(200).set_body_json(body))
.mount(server)
.await;
}
async fn mount_token(server: &MockServer, status: u16, body: serde_json::Value) {
Mock::given(method("POST"))
.and(path("/v1/token"))
.respond_with(ResponseTemplate::new(status).set_body_json(body))
.mount(server)
.await;
}
#[tokio::test]
async fn start_parses_device_authorization() {
let server = MockServer::start().await;
mount_device_authorize(
&server,
serde_json::json!({
"device_code": "DC",
"user_code": "WDJB-MJHT",
"verification_uri": "https://x/activate",
"verification_uri_complete": "https://x/activate?user_code=WDJB-MJHT",
"expires_in": 600,
"interval": 5
}),
)
.await;
let d = client(&server)
.start(&["openid", "offline_access"])
.await
.unwrap();
assert_eq!(d.user_code(), "WDJB-MJHT");
assert_eq!(d.verification_uri(), "https://x/activate");
assert_eq!(
d.verification_uri_complete(),
Some("https://x/activate?user_code=WDJB-MJHT")
);
assert_eq!(d.expires_in(), 600);
assert_eq!(d.interval(), 5);
assert!(!format!("{d:?}").contains("WDJB-MJHT"));
}
#[tokio::test]
async fn start_defaults_interval_when_absent() {
let server = MockServer::start().await;
mount_device_authorize(
&server,
serde_json::json!({
"device_code": "DC",
"user_code": "U",
"verification_uri": "https://x",
"expires_in": 600
}),
)
.await;
let d = client(&server).start(&["openid"]).await.unwrap();
assert_eq!(d.interval(), 5);
}
#[tokio::test]
async fn poll_ready_returns_tokens() {
let server = MockServer::start().await;
mount_device_authorize(
&server,
serde_json::json!({
"device_code": "DC", "user_code": "U",
"verification_uri": "https://x", "expires_in": 600, "interval": 5
}),
)
.await;
mount_token(
&server,
200,
serde_json::json!({
"access_token": "AT",
"token_type": "Bearer",
"refresh_token": "RT",
"expires_in": 3600
}),
)
.await;
let c = client(&server);
let authz = c.start(&["openid"]).await.unwrap();
let t = c.poll(&authz, Some(Duration::from_secs(5))).await.unwrap();
assert_eq!(t.access_token, "AT");
assert_eq!(t.refresh_token.as_deref(), Some("RT"));
assert_eq!(t.expires_in, 3600);
}
#[tokio::test]
async fn poll_expired_maps_to_error() {
let server = MockServer::start().await;
mount_device_authorize(
&server,
serde_json::json!({
"device_code": "DC", "user_code": "U",
"verification_uri": "https://x", "expires_in": 600, "interval": 5
}),
)
.await;
mount_token(&server, 400, serde_json::json!({"error": "expired_token"})).await;
let c = client(&server);
let authz = c.start(&["openid"]).await.unwrap();
assert!(matches!(
c.poll(&authz, Some(Duration::from_secs(5))).await,
Err(AuthError::Expired)
));
}
#[tokio::test]
async fn poll_denied_maps_to_error() {
let server = MockServer::start().await;
mount_device_authorize(
&server,
serde_json::json!({
"device_code": "DC", "user_code": "U",
"verification_uri": "https://x", "expires_in": 600, "interval": 5
}),
)
.await;
mount_token(&server, 400, serde_json::json!({"error": "access_denied"})).await;
let c = client(&server);
let authz = c.start(&["openid"]).await.unwrap();
assert!(matches!(
c.poll(&authz, Some(Duration::from_secs(5))).await,
Err(AuthError::Denied)
));
}
#[tokio::test]
async fn device_authorization_round_trips_and_resumes() {
let server = MockServer::start().await;
mount_device_authorize(
&server,
serde_json::json!({
"device_code": "DC", "user_code": "U",
"verification_uri": "https://x", "expires_in": 600, "interval": 5
}),
)
.await;
mount_token(
&server,
200,
serde_json::json!({"access_token": "AT", "token_type": "Bearer", "expires_in": 3600}),
)
.await;
let started = client(&server).start(&["openid"]).await.unwrap();
let json = serde_json::to_string(&started).unwrap();
let resumed: DeviceAuthorization = serde_json::from_str(&json).unwrap();
let fresh = DeviceFlowClient::new(Url::parse(&server.uri()).unwrap(), "test-client");
let t = fresh
.poll(&resumed, Some(Duration::from_secs(5)))
.await
.unwrap();
assert_eq!(t.access_token, "AT");
}
}