use std::{fmt, time::Duration};
use secrecy::{ExposeSecret, SecretString};
use serde::Deserialize;
use url::Url;
use crate::{
AppRegistration, ConfigError, DEFAULT_REQUEST_TIMEOUT, Endpoints, Sleeper, USER_AGENT,
UserAccessToken,
};
pub const DEVICE_GRANT_TYPE: &str = "urn:ietf:params:oauth:grant-type:device_code";
pub const SLOW_DOWN_INCREMENT: Duration = Duration::from_secs(5);
pub const DEFAULT_POLL_INTERVAL: Duration = Duration::from_secs(5);
pub const MAX_TRANSPORT_RETRIES: u32 = 5;
#[derive(Debug, thiserror::Error)]
pub enum DeviceFlowError {
#[error("the login was declined on GitHub")]
AccessDenied,
#[error("the device code expired before the login was approved; start `auth login` again")]
Expired,
#[error("GitHub did not recognise the device code; start `auth login` again")]
IncorrectDeviceCode,
#[error("the published GitHub App is misconfigured for the device flow: GitHub said {code:?}")]
AppMisconfigured { code: String },
#[error("GitHub returned an unrecognised device-flow error: {code:?}")]
Unexpected { code: String },
#[error(
"GitHub returned a verification URL on {origin:?}, which is not the canonical device \
page; refusing to display it"
)]
UntrustedVerificationUri { origin: String },
#[error("GitHub was unreachable")]
Transport(#[source] reqwest::Error),
#[error("GitHub returned {status} for the device-flow {stage}")]
Status { status: u16, stage: &'static str },
#[error("a device-flow {stage} response could not be decoded")]
Decode {
stage: &'static str,
#[source]
source: serde_json::Error,
},
#[error("GitHub returned {value:?} for {what}, which this client cannot use")]
Malformed { what: &'static str, value: String },
#[error(transparent)]
Config(#[from] ConfigError),
}
impl DeviceFlowError {
#[must_use]
pub fn is_retryable(&self) -> bool {
match self {
Self::Transport(_) => true,
Self::Status { status, .. } => (500..600).contains(status),
Self::AccessDenied
| Self::Expired
| Self::IncorrectDeviceCode
| Self::AppMisconfigured { .. }
| Self::Unexpected { .. }
| Self::UntrustedVerificationUri { .. }
| Self::Decode { .. }
| Self::Malformed { .. }
| Self::Config(_) => false,
}
}
#[must_use]
pub fn requires_new_login(&self) -> bool {
matches!(self, Self::Expired | Self::IncorrectDeviceCode)
}
}
fn transport(err: reqwest::Error) -> DeviceFlowError {
DeviceFlowError::Transport(err.without_url())
}
#[derive(Clone)]
pub struct DeviceAuthorization {
device_code: SecretString,
user_code: String,
verification_uri: Url,
expires_in: Duration,
interval: Duration,
}
impl DeviceAuthorization {
#[must_use]
pub fn user_code(&self) -> &str {
&self.user_code
}
#[must_use]
pub fn device_code(&self) -> &SecretString {
&self.device_code
}
#[must_use]
pub fn verification_uri(&self) -> &Url {
&self.verification_uri
}
#[must_use]
pub fn expires_in(&self) -> Duration {
self.expires_in
}
#[must_use]
pub fn interval(&self) -> Duration {
self.interval
}
}
impl fmt::Debug for DeviceAuthorization {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DeviceAuthorization")
.field("device_code", &"[REDACTED]")
.field("user_code", &self.user_code)
.field("verification_uri", &self.verification_uri.as_str())
.field("expires_in_secs", &self.expires_in.as_secs())
.field("interval_secs", &self.interval.as_secs())
.finish()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PollOutcome {
Pending,
SlowDown { interval: Duration },
Approved(UserAccessToken),
}
#[derive(Debug, Clone)]
pub struct DeviceFlow {
http: reqwest::Client,
app: AppRegistration,
endpoints: Endpoints,
}
impl DeviceFlow {
pub fn new(app: AppRegistration, endpoints: Endpoints) -> Result<Self, DeviceFlowError> {
let http = reqwest::Client::builder()
.timeout(DEFAULT_REQUEST_TIMEOUT)
.build()
.map_err(transport)?;
Ok(Self::with_http_client(http, app, endpoints))
}
pub async fn refresh(
&self,
refresh_token: &SecretString,
) -> Result<UserAccessToken, DeviceFlowError> {
let body = form_body(&[
("client_id", self.app.client_id()),
("refresh_token", refresh_token.expose_secret()),
("grant_type", "refresh_token"),
]);
let response = self
.http
.post(self.endpoints.access_token_url())
.header(reqwest::header::ACCEPT, "application/json")
.header(reqwest::header::USER_AGENT, USER_AGENT)
.header(
reqwest::header::CONTENT_TYPE,
"application/x-www-form-urlencoded",
)
.body(body)
.send()
.await
.map_err(transport)?;
let status = response.status();
let bytes = response.bytes().await.map_err(transport)?;
let raw: RawTokenResponse = match serde_json::from_slice(&bytes) {
Ok(raw) => raw,
Err(source) => {
return if status.is_success() {
Err(DeviceFlowError::Decode {
stage: "refresh",
source,
})
} else {
Err(DeviceFlowError::Status {
status: status.as_u16(),
stage: "refresh request",
})
};
}
};
if let Some(code) = raw.error {
return Err(DeviceFlowError::Malformed {
what: "a refresh response",
value: code,
});
}
let Some(access_token) = raw.access_token else {
return Err(DeviceFlowError::Malformed {
what: "a refresh response",
value: "neither an access token nor an error".to_string(),
});
};
Ok(UserAccessToken::from_parts(
SecretString::from(access_token),
raw.token_type.unwrap_or_else(|| "bearer".to_string()),
raw.scope.filter(|s| !s.is_empty()),
)
.with_renewal(
raw.refresh_token.map(SecretString::from),
raw.expires_in,
raw.refresh_token_expires_in,
))
}
#[must_use]
pub fn with_http_client(
http: reqwest::Client,
app: AppRegistration,
endpoints: Endpoints,
) -> Self {
Self {
http,
app,
endpoints,
}
}
#[must_use]
pub fn verification_url(&self) -> Url {
self.endpoints.verification_url()
}
pub async fn start(&self) -> Result<DeviceAuthorization, DeviceFlowError> {
let body = form_body(&[("client_id", self.app.client_id())]);
let response = self
.http
.post(self.endpoints.device_code_url())
.header(reqwest::header::ACCEPT, "application/json")
.header(reqwest::header::USER_AGENT, USER_AGENT)
.header(
reqwest::header::CONTENT_TYPE,
"application/x-www-form-urlencoded",
)
.body(body)
.send()
.await
.map_err(transport)?;
let status = response.status();
let bytes = response.bytes().await.map_err(transport)?;
if !status.is_success() {
return Err(DeviceFlowError::Status {
status: status.as_u16(),
stage: "device code request",
});
}
let raw: RawDeviceCode =
serde_json::from_slice(&bytes).map_err(|source| DeviceFlowError::Decode {
stage: "device code",
source,
})?;
let verification_uri =
Url::parse(&raw.verification_uri).map_err(|_| DeviceFlowError::Malformed {
what: "a verification URL",
value: raw.verification_uri.clone(),
})?;
if verification_uri.origin() != self.endpoints.web_base().origin() {
return Err(DeviceFlowError::UntrustedVerificationUri {
origin: verification_uri.origin().ascii_serialization(),
});
}
let authorization = DeviceAuthorization {
device_code: SecretString::from(raw.device_code),
user_code: raw.user_code,
verification_uri,
expires_in: Duration::from_secs(raw.expires_in.unwrap_or(900)),
interval: raw
.interval
.map_or(DEFAULT_POLL_INTERVAL, Duration::from_secs),
};
tracing::info!(
user_code = %authorization.user_code,
verification_url = %self.verification_url(),
expires_in_secs = authorization.expires_in.as_secs(),
"device login started; approve it on GitHub's own device page"
);
Ok(authorization)
}
pub async fn poll_once(
&self,
authorization: &DeviceAuthorization,
) -> Result<PollOutcome, DeviceFlowError> {
self.poll_once_from(authorization, authorization.interval)
.await
}
async fn poll_once_from(
&self,
authorization: &DeviceAuthorization,
current_interval: Duration,
) -> Result<PollOutcome, DeviceFlowError> {
let body = form_body(&[
("client_id", self.app.client_id()),
("device_code", authorization.device_code.expose_secret()),
("grant_type", DEVICE_GRANT_TYPE),
]);
let response = self
.http
.post(self.endpoints.access_token_url())
.header(reqwest::header::ACCEPT, "application/json")
.header(reqwest::header::USER_AGENT, USER_AGENT)
.header(
reqwest::header::CONTENT_TYPE,
"application/x-www-form-urlencoded",
)
.body(body)
.send()
.await
.map_err(transport)?;
let status = response.status();
let bytes = response.bytes().await.map_err(transport)?;
let mut raw: RawTokenResponse = match serde_json::from_slice(&bytes) {
Ok(raw) => raw,
Err(source) => {
if status.is_success() {
return Err(DeviceFlowError::Decode {
stage: "access token",
source,
});
}
return Err(DeviceFlowError::Status {
status: status.as_u16(),
stage: "access token request",
});
}
};
if let Some(code) = raw.error.take() {
return self.interpret_error(&code, raw.interval, current_interval);
}
if !status.is_success() {
return Err(DeviceFlowError::Status {
status: status.as_u16(),
stage: "access token request",
});
}
let Some(access_token) = raw.access_token.take() else {
return Err(DeviceFlowError::Malformed {
what: "an access token response",
value: "neither an access token nor an error".to_string(),
});
};
let token = UserAccessToken::from_parts(
SecretString::from(access_token),
raw.token_type.unwrap_or_else(|| "bearer".to_string()),
raw.scope.filter(|s| !s.is_empty()),
)
.with_renewal(
raw.refresh_token.map(SecretString::from),
raw.expires_in,
raw.refresh_token_expires_in,
);
tracing::info!(
token_family = token.family(),
user_to_server = token.is_user_to_server(),
"device login approved; the user access token was returned to the caller"
);
Ok(PollOutcome::Approved(token))
}
fn interpret_error(
&self,
code: &str,
advertised_interval: Option<u64>,
current_interval: Duration,
) -> Result<PollOutcome, DeviceFlowError> {
match code {
"authorization_pending" => Ok(PollOutcome::Pending),
"slow_down" => {
let interval = slowed(
current_interval,
advertised_interval.map(Duration::from_secs),
);
tracing::debug!(
from_secs = current_interval.as_secs(),
to_secs = interval.as_secs(),
"GitHub asked us to slow down; lengthening the poll interval"
);
Ok(PollOutcome::SlowDown { interval })
}
"expired_token" => Err(DeviceFlowError::Expired),
"access_denied" => Err(DeviceFlowError::AccessDenied),
"incorrect_device_code" => Err(DeviceFlowError::IncorrectDeviceCode),
"unsupported_grant_type" | "incorrect_client_credentials" | "device_flow_disabled" => {
Err(DeviceFlowError::AppMisconfigured {
code: code.to_string(),
})
}
other => Err(DeviceFlowError::Unexpected {
code: other.to_string(),
}),
}
}
pub async fn complete(
&self,
authorization: &DeviceAuthorization,
sleeper: &dyn Sleeper,
) -> Result<UserAccessToken, DeviceFlowError> {
let mut interval = authorization.interval;
let mut elapsed = Duration::ZERO;
let mut consecutive_retryable = 0_u32;
loop {
sleeper.sleep(interval).await;
elapsed = elapsed.saturating_add(interval);
let outcome = match self.poll_once_from(authorization, interval).await {
Ok(outcome) => {
consecutive_retryable = 0;
outcome
}
Err(err) if err.is_retryable() && consecutive_retryable < MAX_TRANSPORT_RETRIES => {
consecutive_retryable += 1;
tracing::warn!(
attempt = consecutive_retryable,
max_attempts = MAX_TRANSPORT_RETRIES,
error = %err,
"a device-flow poll failed in a way that says nothing about the login; \
the device code is still live, so polling continues"
);
PollOutcome::Pending
}
Err(err) => return Err(err),
};
match outcome {
PollOutcome::Approved(token) => return Ok(token),
PollOutcome::SlowDown { interval: next } => interval = next,
PollOutcome::Pending => {}
}
if elapsed >= authorization.expires_in {
tracing::warn!(
waited_secs = elapsed.as_secs(),
expires_in_secs = authorization.expires_in.as_secs(),
"the device code's own lifetime elapsed before approval"
);
return Err(DeviceFlowError::Expired);
}
}
}
}
fn slowed(current: Duration, advertised: Option<Duration>) -> Duration {
let floor = current.saturating_add(SLOW_DOWN_INCREMENT);
advertised.map_or(floor, |advertised| advertised.max(floor))
}
fn form_body(pairs: &[(&str, &str)]) -> String {
let mut serializer = url::form_urlencoded::Serializer::new(String::new());
for (key, value) in pairs {
serializer.append_pair(key, value);
}
serializer.finish()
}
#[derive(Debug, Deserialize)]
struct RawDeviceCode {
device_code: String,
user_code: String,
verification_uri: String,
#[serde(default)]
expires_in: Option<u64>,
#[serde(default)]
interval: Option<u64>,
}
#[derive(Deserialize)]
struct RawTokenResponse {
#[serde(default)]
access_token: Option<String>,
#[serde(default)]
token_type: Option<String>,
#[serde(default)]
scope: Option<String>,
#[serde(default)]
refresh_token: Option<String>,
#[serde(default)]
expires_in: Option<u64>,
#[serde(default)]
refresh_token_expires_in: Option<u64>,
#[serde(default)]
error: Option<String>,
#[serde(default)]
interval: Option<u64>,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::testing::{
FIXTURE_DEVICE_CODE, FIXTURE_TOKEN, FIXTURE_USER_CODE, RecordingSleeper, Script,
device_code_body, error_body, token_body,
};
use serde_json::json;
use wiremock::{
Mock, MockServer, ResponseTemplate,
matchers::{body_string_contains, header, method, path},
};
fn app() -> AppRegistration {
AppRegistration::new("Iv23liTESTCLIENTID", "runner-manager").unwrap()
}
fn flow(server: &MockServer) -> DeviceFlow {
DeviceFlow::new(app(), Endpoints::for_test_server(&server.uri()).unwrap()).unwrap()
}
async fn mount_start(server: &MockServer) {
Mock::given(method("POST"))
.and(path("/login/device/code"))
.respond_with(ResponseTemplate::new(200).set_body_json(device_code_body(
&server.uri(),
5,
900,
)))
.mount(server)
.await;
}
async fn mount_token(server: &MockServer, responses: Vec<ResponseTemplate>) {
Mock::given(method("POST"))
.and(path("/login/oauth/access_token"))
.respond_with(Script::new(responses))
.mount(server)
.await;
}
#[tokio::test]
async fn a_device_flow_round_trip_against_fixtures_succeeds() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![ResponseTemplate::new(200).set_body_json(token_body())],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.expect("device code");
assert_eq!(authorization.user_code(), FIXTURE_USER_CODE);
assert_eq!(
authorization.device_code().expose_secret(),
FIXTURE_DEVICE_CODE
);
assert_eq!(authorization.interval(), Duration::from_secs(5));
assert_eq!(authorization.expires_in(), Duration::from_secs(900));
let sleeper = RecordingSleeper::default();
let token = flow
.complete(&authorization, &sleeper)
.await
.expect("approved");
assert_eq!(token.secret().expose_secret(), FIXTURE_TOKEN);
assert_eq!(token.token_type(), "bearer");
assert_eq!(token.family(), "ghu_");
assert!(
token.is_user_to_server(),
"the published App issues user-to-server tokens; anything else means the \
registration is not the one this product authenticates as"
);
assert_eq!(sleeper.recorded(), vec![Duration::from_secs(5)]);
}
#[tokio::test]
async fn the_start_request_carries_the_public_client_id_and_no_secret() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/login/device/code"))
.and(header("accept", "application/json"))
.and(header("content-type", "application/x-www-form-urlencoded"))
.and(body_string_contains("client_id=Iv23liTESTCLIENTID"))
.respond_with(ResponseTemplate::new(200).set_body_json(device_code_body(
&server.uri(),
5,
900,
)))
.expect(1)
.mount(&server)
.await;
flow(&server).start().await.expect("device code");
let sent = server.received_requests().await.unwrap();
let body = String::from_utf8(sent[0].body.clone()).unwrap();
assert_eq!(
body, "client_id=Iv23liTESTCLIENTID",
"the start request is the client id and nothing else: no secret, no scope, \
no redirect URI"
);
}
#[tokio::test]
async fn the_device_code_travels_in_the_body_and_never_in_the_url() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![ResponseTemplate::new(200).set_body_json(token_body())],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
flow.complete(&authorization, &RecordingSleeper::default())
.await
.unwrap();
let sent = server.received_requests().await.unwrap();
let poll = sent
.iter()
.find(|r| r.url.path() == "/login/oauth/access_token")
.expect("the poll happened");
assert!(
!poll.url.as_str().contains(FIXTURE_DEVICE_CODE),
"a query string reaches every proxy and access log: {}",
poll.url
);
let body = String::from_utf8(poll.body.clone()).unwrap();
assert!(body.contains(&format!("device_code={FIXTURE_DEVICE_CODE}")));
assert!(body.contains("grant_type=urn%3Aietf%3Aparams%3Aoauth%3Agrant-type%3Adevice_code"));
}
#[tokio::test]
async fn authorization_pending_keeps_polling_at_the_unchanged_interval() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![
ResponseTemplate::new(200).set_body_json(error_body("authorization_pending", None)),
ResponseTemplate::new(200).set_body_json(error_body("authorization_pending", None)),
ResponseTemplate::new(200).set_body_json(error_body("authorization_pending", None)),
ResponseTemplate::new(200).set_body_json(token_body()),
],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
assert_eq!(
flow.poll_once(&authorization).await.unwrap(),
PollOutcome::Pending,
"pending is an outcome, not a failure"
);
let sleeper = RecordingSleeper::default();
flow.complete(&authorization, &sleeper).await.unwrap();
assert_eq!(
sleeper.recorded(),
vec![
Duration::from_secs(5),
Duration::from_secs(5),
Duration::from_secs(5)
],
"authorization_pending must not change the interval"
);
}
#[tokio::test]
async fn slow_down_demonstrably_increases_the_poll_interval() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![
ResponseTemplate::new(200).set_body_json(error_body("authorization_pending", None)),
ResponseTemplate::new(200).set_body_json(error_body("slow_down", Some(10))),
ResponseTemplate::new(200).set_body_json(error_body("authorization_pending", None)),
ResponseTemplate::new(200).set_body_json(error_body("slow_down", Some(15))),
ResponseTemplate::new(200).set_body_json(token_body()),
],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let sleeper = RecordingSleeper::default();
flow.complete(&authorization, &sleeper).await.unwrap();
let waited = sleeper.recorded();
assert_eq!(
waited,
vec![
Duration::from_secs(5), Duration::from_secs(5), Duration::from_secs(10), Duration::from_secs(10),
Duration::from_secs(15), ],
"every slow_down must lengthen the interval used from then on"
);
let mut increases = 0;
for pair in waited.windows(2) {
assert!(
pair[1] >= pair[0],
"the interval must never shrink: {pair:?}"
);
if pair[1] > pair[0] {
increases += 1;
}
}
assert_eq!(increases, 2, "two slow_downs, two increases");
}
#[tokio::test]
async fn a_slow_down_without_an_advertised_interval_still_adds_five_seconds() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![
ResponseTemplate::new(200).set_body_json(error_body("slow_down", None)),
ResponseTemplate::new(200).set_body_json(token_body()),
],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
assert_eq!(
flow.poll_once(&authorization).await.unwrap(),
PollOutcome::SlowDown {
interval: Duration::from_secs(10)
},
"RFC 8628 makes the +5s increase mandatory even with no interval in the body"
);
}
#[test]
fn the_slow_down_interval_never_shrinks_whatever_the_server_advertises() {
assert_eq!(
slowed(Duration::from_secs(5), None),
Duration::from_secs(10)
);
assert_eq!(
slowed(Duration::from_secs(5), Some(Duration::from_secs(10))),
Duration::from_secs(10),
"GitHub advertises exactly the RFC floor, so the two agree"
);
assert_eq!(
slowed(Duration::from_secs(5), Some(Duration::from_secs(30))),
Duration::from_secs(30),
"a server asking for more than the floor gets it"
);
assert_eq!(
slowed(Duration::from_secs(5), Some(Duration::from_secs(1))),
Duration::from_secs(10),
"a server asking for LESS must not be able to speed us up: slow_down means slow down"
);
}
#[tokio::test]
async fn expired_token_is_terminal_and_asks_for_a_whole_new_login() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![ResponseTemplate::new(200).set_body_json(error_body("expired_token", None))],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let err = flow
.complete(&authorization, &RecordingSleeper::default())
.await
.expect_err("expired");
assert!(matches!(err, DeviceFlowError::Expired), "{err:?}");
assert!(!err.is_retryable());
assert!(err.requires_new_login());
assert!(err.to_string().contains("auth login"));
}
#[tokio::test]
async fn access_denied_is_terminal_and_is_not_an_error_to_retry() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![ResponseTemplate::new(200).set_body_json(error_body("access_denied", None))],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let err = flow
.complete(&authorization, &RecordingSleeper::default())
.await
.expect_err("declined");
assert!(matches!(err, DeviceFlowError::AccessDenied), "{err:?}");
assert!(!err.is_retryable(), "the user said no; do not ask again");
assert!(
!err.requires_new_login(),
"a refusal is not an expiry: `auth login` again is the operator's choice, \
not this error's instruction"
);
}
#[tokio::test]
async fn a_gateway_blip_mid_poll_does_not_abort_the_login() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![
ResponseTemplate::new(502).set_body_string("<html>Bad Gateway</html>"),
ResponseTemplate::new(200).set_body_json(error_body("authorization_pending", None)),
ResponseTemplate::new(503).set_body_string("<html>Service Unavailable</html>"),
ResponseTemplate::new(200).set_body_json(token_body()),
],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let sleeper = RecordingSleeper::default();
let token = flow
.complete(&authorization, &sleeper)
.await
.expect("two blips must not destroy a login the user is about to approve");
assert_eq!(token.secret().expose_secret(), FIXTURE_TOKEN);
assert_eq!(
sleeper.recorded().len(),
4,
"each absorbed failure costs one more scheduled poll and nothing else"
);
}
#[tokio::test]
async fn the_transport_retry_is_bounded_rather_than_endless() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![ResponseTemplate::new(503).set_body_string("down")],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let sleeper = RecordingSleeper::default();
let err = flow
.complete(&authorization, &sleeper)
.await
.expect_err("a persistently unreachable GitHub is still a failure");
assert!(
matches!(err, DeviceFlowError::Status { status: 503, .. }),
"{err:?}"
);
assert_eq!(
sleeper.recorded().len(),
MAX_TRANSPORT_RETRIES as usize + 1,
"one initial poll plus exactly {MAX_TRANSPORT_RETRIES} absorbed retries"
);
}
#[tokio::test]
async fn a_transport_failure_is_absorbed_and_bounded_like_a_gateway_error() {
let server = MockServer::start().await;
mount_start(&server).await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let dead_port = {
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("an ephemeral port");
listener.local_addr().expect("a bound address").port()
};
let unreachable = DeviceFlow::new(
app(),
Endpoints::for_test_server(&format!("http://127.0.0.1:{dead_port}")).unwrap(),
)
.unwrap();
let sleeper = RecordingSleeper::default();
let err = unreachable
.complete(&authorization, &sleeper)
.await
.expect_err("nothing is listening");
assert!(matches!(err, DeviceFlowError::Transport(_)), "{err:?}");
assert!(
err.is_retryable(),
"the device code is still live, so `f1` may present this same login again"
);
assert_eq!(
sleeper.recorded().len(),
MAX_TRANSPORT_RETRIES as usize + 1,
"a dropped connection is absorbed like any other blip, and bounded the same way"
);
}
#[tokio::test]
async fn a_proxy_error_page_is_reported_as_its_status_not_as_a_decode_failure() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![
ResponseTemplate::new(502)
.insert_header("content-type", "text/html")
.set_body_string("<html><body>502 Bad Gateway</body></html>"),
],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let err = flow
.poll_once(&authorization)
.await
.expect_err("a gateway page is not a token response");
match err {
DeviceFlowError::Status { status, stage } => {
assert_eq!(status, 502);
assert_eq!(stage, "access token request");
}
other => panic!("expected the status, got {other:?}"),
}
}
#[tokio::test]
async fn a_200_that_is_not_the_protocol_is_still_a_decode_failure() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![ResponseTemplate::new(200).set_body_string("not json at all")],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let err = flow.poll_once(&authorization).await.expect_err("garbage");
assert!(matches!(err, DeviceFlowError::Decode { .. }), "{err:?}");
assert!(
!err.is_retryable(),
"a 200 that is not the protocol is a protocol violation, not a blip"
);
}
#[tokio::test]
async fn the_terminal_states_are_never_retried() {
for code in [
"access_denied",
"expired_token",
"incorrect_device_code",
"device_flow_disabled",
"a_code_this_client_has_never_heard_of",
] {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![ResponseTemplate::new(200).set_body_json(error_body(code, None))],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let sleeper = RecordingSleeper::default();
let err = flow
.complete(&authorization, &sleeper)
.await
.expect_err("terminal");
assert!(!err.is_retryable(), "{code} must stay terminal: {err:?}");
assert_eq!(
sleeper.recorded().len(),
1,
"{code} must be answered on the first poll and never polled again"
);
}
}
#[tokio::test]
async fn the_four_documented_errors_produce_four_distinct_outcomes() {
let server = MockServer::start().await;
mount_start(&server).await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let mut described = Vec::new();
for code in [
"authorization_pending",
"slow_down",
"expired_token",
"access_denied",
] {
let scoped = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/login/oauth/access_token"))
.respond_with(ResponseTemplate::new(200).set_body_json(error_body(code, None)))
.mount(&scoped)
.await;
let scoped_flow =
DeviceFlow::new(app(), Endpoints::for_test_server(&scoped.uri()).unwrap()).unwrap();
described.push(match scoped_flow.poll_once(&authorization).await {
Ok(PollOutcome::Pending) => "pending".to_string(),
Ok(PollOutcome::SlowDown { interval }) => {
format!("slow_down->{}s", interval.as_secs())
}
Ok(PollOutcome::Approved(_)) => "approved".to_string(),
Err(err) => format!("error:{err:?}"),
});
}
assert_eq!(
described,
vec![
"pending",
"slow_down->10s",
"error:Expired",
"error:AccessDenied",
]
);
let mut unique = described.clone();
unique.sort();
unique.dedup();
assert_eq!(unique.len(), 4, "four codes must not collapse into fewer");
}
#[tokio::test]
async fn an_unrecognised_error_is_reported_as_itself_rather_than_guessed_at() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![
ResponseTemplate::new(200).set_body_json(error_body("device_flow_disabled", None)),
],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let err = flow.poll_once(&authorization).await.expect_err("disabled");
match err {
DeviceFlowError::AppMisconfigured { code } => {
assert_eq!(code, "device_flow_disabled");
}
other => panic!("expected a registration error, got {other:?}"),
}
}
#[tokio::test]
async fn the_local_deadline_backstops_a_server_that_never_says_expired() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/login/device/code"))
.respond_with(ResponseTemplate::new(200).set_body_json(device_code_body(
&server.uri(),
5,
20,
)))
.mount(&server)
.await;
mount_token(
&server,
vec![
ResponseTemplate::new(200).set_body_json(error_body("authorization_pending", None)),
],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let sleeper = RecordingSleeper::default();
let err = flow
.complete(&authorization, &sleeper)
.await
.expect_err("the device code's own lifetime ran out");
assert!(matches!(err, DeviceFlowError::Expired), "{err:?}");
assert_eq!(
sleeper.recorded().len(),
4,
"the loop stops at the advertised lifetime rather than polling forever"
);
}
#[tokio::test]
async fn a_verification_url_on_another_origin_is_refused_rather_than_displayed() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/login/device/code"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"device_code": FIXTURE_DEVICE_CODE,
"user_code": FIXTURE_USER_CODE,
"verification_uri": "https://github.com.evil.example/login/device",
"expires_in": 900,
"interval": 5
})))
.mount(&server)
.await;
let err = flow(&server).start().await.expect_err("wrong origin");
match err {
DeviceFlowError::UntrustedVerificationUri { origin } => {
assert!(origin.contains("evil.example"), "{origin}");
}
other => panic!("expected the phishing control to fire, got {other:?}"),
}
}
#[test]
fn the_printed_verification_url_is_the_compiled_in_canonical_one() {
let production = DeviceFlow::new(app(), Endpoints::production()).unwrap();
assert_eq!(
production.verification_url().as_str(),
"https://github.com/login/device",
"the tool prints this and never proxies, embeds, or imitates the approval page"
);
}
#[tokio::test]
async fn neither_code_nor_token_is_rendered_by_debug_or_display() {
let server = MockServer::start().await;
mount_start(&server).await;
mount_token(
&server,
vec![ResponseTemplate::new(200).set_body_json(token_body())],
)
.await;
let flow = flow(&server);
let authorization = flow.start().await.unwrap();
let rendered = format!("{authorization:?}");
assert!(
!rendered.contains(FIXTURE_DEVICE_CODE),
"the device code is never shown: {rendered}"
);
assert!(rendered.contains("[REDACTED]"));
assert!(
rendered.contains(FIXTURE_USER_CODE),
"the user code is displayed by design, so it stays legible: {rendered}"
);
let token = flow
.complete(&authorization, &RecordingSleeper::default())
.await
.unwrap();
assert!(!format!("{token:?}").contains(FIXTURE_TOKEN));
assert!(!format!("{flow:?}").contains(FIXTURE_TOKEN));
for err in [
DeviceFlowError::AccessDenied,
DeviceFlowError::Expired,
DeviceFlowError::IncorrectDeviceCode,
DeviceFlowError::AppMisconfigured {
code: "device_flow_disabled".to_string(),
},
DeviceFlowError::Unexpected {
code: "??".to_string(),
},
] {
let text = format!("{err} / {err:?}");
assert!(!text.contains(FIXTURE_DEVICE_CODE), "{text}");
assert!(!text.contains(FIXTURE_TOKEN), "{text}");
}
}
#[test]
fn the_token_family_is_diagnostic_and_exposes_nothing_else() {
let token = UserAccessToken::new(SecretString::from("ghu_abcdefghijklmnop"));
assert_eq!(token.family(), "ghu_");
assert!(token.is_user_to_server());
let oauth = UserAccessToken::new(SecretString::from("gho_abcdefghijklmnop"));
assert_eq!(oauth.family(), "gho_");
assert!(
!oauth.is_user_to_server(),
"a `gho_` token means this is not the published App's user-to-server credential"
);
let odd = UserAccessToken::new(SecretString::from("no-underscore-here"));
assert_eq!(
odd.family(),
"",
"never guess, and never return a prefix of the token"
);
}
#[test]
fn the_form_body_encodes_exactly_what_both_spikes_sent() {
assert_eq!(
form_body(&[("client_id", "Iv1"), ("grant_type", DEVICE_GRANT_TYPE)]),
"client_id=Iv1&grant_type=urn%3Aietf%3Aparams%3Aoauth%3Agrant-type%3Adevice_code"
);
}
}