use std::collections::BTreeSet;
use oauth2::{
ClientId, DeviceAuthorizationUrl, DeviceCodeErrorResponse, DeviceCodeErrorResponseType,
HttpClientError, RequestTokenError, Scope, StandardDeviceAuthorizationResponse,
StandardErrorResponse, TokenUrl,
basic::{BasicClient, BasicErrorResponseType},
};
use tokio_util::sync::CancellationToken;
use url::Url;
use crate::configuration::login::{LoginResponse, oauth_http_client, resolve_scopes};
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum DeviceLoginError {
#[error(transparent)]
ReqwestClient(#[from] oauth2::reqwest::Error),
#[error("Failed to request a device authorization code: {0}")]
DeviceAuthorization(
#[from]
RequestTokenError<
HttpClientError<oauth2::reqwest::Error>,
StandardErrorResponse<BasicErrorResponseType>,
>,
),
#[error("Failed to exchange the device code for an access token: {0}")]
RequestToken(
#[from] RequestTokenError<HttpClientError<oauth2::reqwest::Error>, DeviceCodeErrorResponse>,
),
#[error("The device authorization login was cancelled")]
Cancelled,
#[error(
"The OAuth issuer does not advertise support for the device authorization grant. It must \
publish a `device_authorization_endpoint` in its discovery document."
)]
NotSupported,
}
impl DeviceLoginError {
pub(crate) fn allows_pkce_fallback(&self) -> bool {
match self {
Self::Cancelled => false,
Self::RequestToken(RequestTokenError::ServerResponse(response)) => {
!matches!(response.error(), DeviceCodeErrorResponseType::AccessDenied)
}
_ => true,
}
}
}
pub(crate) enum DevicePrompt {
User,
#[cfg(test)]
Test(tokio::sync::mpsc::UnboundedSender<String>),
}
impl DevicePrompt {
async fn show(self, details: &StandardDeviceAuthorizationResponse) {
match self {
Self::User => {
let verification_uri = details.verification_uri();
println!(
"Login to QCS by going to {verification_uri} and entering the code: {user_code}",
user_code = details.user_code().secret(),
);
let browser_uri = details
.verification_uri_complete()
.map_or_else(|| verification_uri.to_string(), |uri| uri.secret().clone());
_ = tokio::task::spawn_blocking(move || webbrowser::open(&browser_uri)).await;
}
#[cfg(test)]
Self::Test(prompted_tx) => {
_ = prompted_tx.send(details.user_code().secret().clone());
}
}
}
}
pub(crate) struct DeviceLoginRequest {
pub(crate) client_id: String,
pub(crate) token_endpoint: Url,
pub(crate) device_authorization_endpoint: Url,
pub(crate) scopes: Option<BTreeSet<String>>,
pub(crate) advertised_scopes: Option<BTreeSet<String>>,
pub(crate) prompt: DevicePrompt,
}
pub(crate) async fn device_login(
cancel_token: CancellationToken,
request: DeviceLoginRequest,
) -> Result<LoginResponse, DeviceLoginError> {
let DeviceLoginRequest {
client_id,
token_endpoint,
device_authorization_endpoint,
scopes,
advertised_scopes,
prompt,
} = request;
let client = BasicClient::new(ClientId::new(client_id))
.set_token_uri(TokenUrl::from_url(token_endpoint))
.set_device_authorization_url(DeviceAuthorizationUrl::from_url(
device_authorization_endpoint,
));
let scopes = resolve_scopes(scopes, advertised_scopes);
let http_client = oauth_http_client()?;
let details: StandardDeviceAuthorizationResponse = client
.exchange_device_code()
.add_scopes(scopes.into_iter().map(Scope::new))
.request_async(&http_client)
.await?;
prompt.show(&details).await;
cancel_token
.run_until_cancelled(client.exchange_device_access_token(&details).request_async(
&http_client,
tokio::time::sleep,
None,
))
.await
.ok_or(DeviceLoginError::Cancelled)?
.map_err(DeviceLoginError::RequestToken)
}
#[cfg(test)]
mod tests {
use httpmock::prelude::*;
use oauth2::TokenResponse;
use oauth2_test_server::{Client as TestClient, IssuerConfig, OAuthTestServer};
use rstest::rstest;
use serde_json::json;
use tokio_util::sync::CancellationToken;
use crate::configuration::{
login::tests::default_scope_string, oidc::DEVICE_CODE_GRANT_TYPE,
secrets::SecretAccessToken, tokens::insecure_validate_token_exp,
};
use super::*;
const DEVICE_AUTHORIZE_PATH: &str = "/v1/device/authorize";
const DEVICE_CODE_EXPIRES_IN_SECS: u64 = 15;
async fn start_device_server() -> (OAuthTestServer, TestClient) {
let server = OAuthTestServer::start_with_config(IssuerConfig::default()).await;
let client = server
.register_client(json!({
"scope": default_scope_string(),
"grant_types": [DEVICE_CODE_GRANT_TYPE],
"client_name": "device-flow-test",
}))
.await;
(server, client)
}
async fn create_device_code(server: &OAuthTestServer, client_id: &str) -> (String, String) {
let response = server
.http
.post(format!("{}/device/code", server.issuer()))
.form(&[("client_id", client_id), ("scope", &default_scope_string())])
.send()
.await
.expect("device code request should succeed");
let body: serde_json::Value = response
.json()
.await
.expect("device code response should be JSON");
let field = |name: &str| {
body.get(name)
.and_then(serde_json::Value::as_str)
.unwrap_or_else(|| panic!("device code response should have a `{name}`: {body}"))
.to_string()
};
(field("device_code"), field("user_code"))
}
#[tokio::test(flavor = "multi_thread")]
async fn test_device_login() {
let (oauth_server, client) = start_device_server().await;
let (device_code, user_code) = create_device_code(&oauth_server, &client.client_id).await;
let mock_server = MockServer::start_async().await;
let device_authorize_mock = mock_server
.mock_async(|when, then| {
when.method(POST)
.path(DEVICE_AUTHORIZE_PATH)
.form_urlencoded_tuple("client_id", &client.client_id)
.form_urlencoded_tuple("scope", default_scope_string());
then.status(200).json_body(json!({
"device_code": device_code,
"user_code": user_code,
"verification_uri": format!("{}/device", oauth_server.issuer()),
"expires_in": DEVICE_CODE_EXPIRES_IN_SECS,
"interval": 1,
}));
})
.await;
let (prompted_tx, mut prompted_rx) = tokio::sync::mpsc::unbounded_channel();
let request = DeviceLoginRequest {
client_id: client.client_id.clone(),
token_endpoint: format!("{}/device/token", oauth_server.issuer())
.parse()
.unwrap(),
device_authorization_endpoint: mock_server.url(DEVICE_AUTHORIZE_PATH).parse().unwrap(),
scopes: None,
advertised_scopes: None,
prompt: DevicePrompt::Test(prompted_tx),
};
let (response, prompted_user_code) =
tokio::join!(device_login(CancellationToken::new(), request), async {
let prompted = prompted_rx
.recv()
.await
.expect("the login should prompt the user");
oauth_server
.approve_device_code(&device_code, "device-test-user")
.await;
prompted
});
let response = response.expect("device authorization login should succeed");
device_authorize_mock.assert_async().await;
assert_eq!(prompted_user_code, user_code);
let access_token = SecretAccessToken::from(response.access_token().secret().clone());
insecure_validate_token_exp(&access_token).expect("access token should be valid");
assert!(
response.refresh_token().is_some(),
"the device token endpoint should have returned a refresh token"
);
}
fn token_error(error: DeviceCodeErrorResponseType) -> DeviceLoginError {
DeviceLoginError::RequestToken(RequestTokenError::ServerResponse(
DeviceCodeErrorResponse::new(error, None, None),
))
}
#[rstest]
#[case(DeviceLoginError::Cancelled, false)]
#[case(token_error(DeviceCodeErrorResponseType::AccessDenied), false)]
#[case(token_error(DeviceCodeErrorResponseType::ExpiredToken), true)]
#[case(
token_error(DeviceCodeErrorResponseType::Basic(
BasicErrorResponseType::UnauthorizedClient
)),
true
)]
fn test_allows_pkce_fallback(#[case] error: DeviceLoginError, #[case] expected: bool) {
assert_eq!(error.allows_pkce_fallback(), expected, "error: {error}");
}
}