| 1 | //! One RFC 8628 device-authorization polling loop, shared by every Codewhale |
| 2 | //! device-code flow (xAI/Grok device login, Codewhale account login). |
| 3 | //! |
| 4 | //! Ported from pi (<https://github.com/badlogic/pi-mono>), MIT licensed, |
| 5 | //! Copyright (c) 2025 Mario Zechner — see |
| 6 | //! `packages/ai/src/auth/oauth/device-code.ts` for the original |
| 7 | //! `pollOAuthDeviceCodeFlow`. The accumulated behaviours carried over from it: |
| 8 | //! |
| 9 | //! * the RFC 8628 §3.2 default of 5 seconds when the server omits `interval`; |
| 10 | //! * `slow_down` handling that **prefers a server-supplied interval** over the |
| 11 | //! client-tracked one. Trusting only the client-tracked value lets WSL/VM |
| 12 | //! clock drift poll early forever; RFC 8628 §3.5's +5s step is the fallback; |
| 13 | //! * a hard deadline derived from `expires_in`, never slept past even after |
| 14 | //! `slow_down` backoff; |
| 15 | //! * a distinct timeout message when at least one `slow_down` was seen, so the |
| 16 | //! clock-drift case is diagnosable instead of looking like a plain timeout. |
| 17 | //! |
| 18 | //! The loop is generic over the poll result and does no I/O of its own: the |
| 19 | //! caller supplies the poll and the sleep. Nothing here ever holds, formats, or |
| 20 | //! logs a token — `T` is opaque to this module and is never `Debug`-printed. |
| 21 | |
| 22 | use std::time::{Duration, Instant}; |
| 23 | |
| 24 | use anyhow::{Result, bail}; |
| 25 | |
| 26 | /// RFC 8628 §3.2: when the authorization server omits `interval`, clients must |
| 27 | /// poll no faster than every 5 seconds. |
| 28 | pub const DEFAULT_POLL_INTERVAL_SECS: u64 = 5; |
| 29 | /// RFC 8628 §3.5: `slow_down` increases the polling interval by 5 seconds. |
| 30 | pub const SLOW_DOWN_STEP_SECS: u64 = 5; |
| 31 | /// Never poll faster than once a second, whatever the server asks for. |
| 32 | const MINIMUM_INTERVAL: Duration = Duration::from_secs(1); |
| 33 | /// Longest lifetime a run honours. `expires_in` is untrusted server input, and |
| 34 | /// `Instant + Duration` panics when the sum is unrepresentable, so a hostile |
| 35 | /// or corrupt value is clamped here rather than at every caller. No real |
| 36 | /// device grant lives anywhere near a day. |
| 37 | const MAX_LIFETIME: Duration = Duration::from_secs(24 * 60 * 60); |
| 38 | |
| 39 | /// What one poll of the token endpoint told us. |
| 40 | /// |
| 41 | /// A terminal failure is reported by returning `Err` from the poll closure, so |
| 42 | /// each provider keeps its own error text. |
| 43 | pub enum DevicePollOutcome<T> { |
| 44 | /// The user approved; `T` is the provider's parsed token material. |
| 45 | Complete(T), |
| 46 | /// `authorization_pending` — keep the current interval. |
| 47 | Pending, |
| 48 | /// `slow_down` — back off. `interval_seconds` is the server's new minimum |
| 49 | /// when it supplied one (preferred over the client-tracked interval). |
| 50 | SlowDown { interval_seconds: Option<u64> }, |
| 51 | } |
| 52 | |
| 53 | /// A configured device-code polling run. Build one, then [`DeviceCodePoll::run`]. |
| 54 | pub struct DeviceCodePoll { |
| 55 | interval: Duration, |
| 56 | max_interval: Option<Duration>, |
| 57 | lifetime: Duration, |
| 58 | wait_before_first_poll: bool, |
| 59 | timeout_message: String, |
| 60 | slow_down_timeout_message: Option<String>, |
| 61 | } |
| 62 | |
| 63 | impl DeviceCodePoll { |
| 64 | /// Start a run that gives up after `lifetime` with `timeout_message`. |
| 65 | /// |
| 66 | /// The interval starts at the RFC 8628 default of 5 seconds; callers pass |
| 67 | /// the server's `interval` through [`DeviceCodePoll::interval_seconds`]. |
| 68 | #[must_use] |
| 69 | pub fn new(lifetime: Duration, timeout_message: impl Into<String>) -> Self { |
| 70 | Self { |
| 71 | interval: Duration::from_secs(DEFAULT_POLL_INTERVAL_SECS), |
| 72 | max_interval: None, |
| 73 | lifetime: lifetime.min(MAX_LIFETIME), |
| 74 | wait_before_first_poll: false, |
| 75 | timeout_message: timeout_message.into(), |
| 76 | slow_down_timeout_message: None, |
| 77 | } |
| 78 | } |
| 79 | |
| 80 | /// Apply the server-advertised `interval`. `None` (or a zero/absent value, |
| 81 | /// which RFC 8628 permits) keeps the 5-second default. |
| 82 | #[must_use] |
| 83 | pub fn interval_seconds(mut self, seconds: Option<u64>) -> Self { |
| 84 | if let Some(seconds) = seconds.filter(|seconds| *seconds > 0) { |
| 85 | self.interval = self.clamp_interval(Duration::from_secs(seconds)); |
| 86 | } |
| 87 | self |
| 88 | } |
| 89 | |
| 90 | /// Cap the interval, including after `slow_down` backoff. |
| 91 | #[must_use] |
| 92 | pub fn max_interval_seconds(mut self, seconds: u64) -> Self { |
| 93 | self.max_interval = Some(Duration::from_secs(seconds.max(1))); |
| 94 | self.interval = self.clamp_interval(self.interval); |
| 95 | self |
| 96 | } |
| 97 | |
| 98 | /// Sleep one interval before the first poll. |
| 99 | /// |
| 100 | /// Device-code endpoints that answer `authorization_pending` (xAI) want |
| 101 | /// this; endpoints whose first response is already meaningful (the |
| 102 | /// Codewhale account service, which returns HTTP 202 while pending) poll |
| 103 | /// immediately and sleep afterwards. |
| 104 | #[must_use] |
| 105 | pub fn wait_before_first_poll(mut self, wait: bool) -> Self { |
| 106 | self.wait_before_first_poll = wait; |
| 107 | self |
| 108 | } |
| 109 | |
| 110 | /// Message used instead of the plain timeout message when the run saw at |
| 111 | /// least one `slow_down`. This is the WSL/VM clock-drift tell. |
| 112 | #[must_use] |
| 113 | pub fn slow_down_timeout_message(mut self, message: impl Into<String>) -> Self { |
| 114 | self.slow_down_timeout_message = Some(message.into()); |
| 115 | self |
| 116 | } |
| 117 | |
| 118 | fn clamp_interval(&self, interval: Duration) -> Duration { |
| 119 | let interval = interval.max(MINIMUM_INTERVAL); |
| 120 | match self.max_interval { |
| 121 | Some(max) => interval.min(max), |
| 122 | None => interval, |
| 123 | } |
| 124 | } |
| 125 | |
| 126 | /// Poll until the flow completes, fails, or the deadline passes. |
| 127 | /// |
| 128 | /// `sleep` is injected so tests never wait in real time. `poll` returns |
| 129 | /// `Err` for any terminal failure (denied, expired, transport error). |
| 130 | pub fn run<T, S, P>(self, mut sleep: S, mut poll: P) -> Result<T> |
| 131 | where |
| 132 | S: FnMut(Duration), |
| 133 | P: FnMut() -> Result<DevicePollOutcome<T>>, |
| 134 | { |
| 135 | let deadline = Instant::now() + self.lifetime; |
| 136 | let mut interval = self.interval; |
| 137 | let mut saw_slow_down = false; |
| 138 | |
| 139 | if self.wait_before_first_poll { |
| 140 | let remaining = deadline.saturating_duration_since(Instant::now()); |
| 141 | if remaining.is_zero() { |
| 142 | return Err(self.timed_out(saw_slow_down)); |
| 143 | } |
| 144 | sleep(interval.min(remaining)); |
| 145 | } |
| 146 | |
| 147 | while Instant::now() < deadline { |
| 148 | match poll()? { |
| 149 | DevicePollOutcome::Complete(value) => return Ok(value), |
| 150 | DevicePollOutcome::Pending => {} |
| 151 | DevicePollOutcome::SlowDown { interval_seconds } => { |
| 152 | saw_slow_down = true; |
| 153 | // Prefer the server's new minimum when it gave one: a |
| 154 | // purely client-tracked interval polls early forever when |
| 155 | // the clock drifts (WSL, suspended VMs). |
| 156 | interval = match interval_seconds.filter(|seconds| *seconds > 0) { |
| 157 | Some(seconds) => self.clamp_interval(Duration::from_secs(seconds)), |
| 158 | None => self.clamp_interval( |
| 159 | interval.saturating_add(Duration::from_secs(SLOW_DOWN_STEP_SECS)), |
| 160 | ), |
| 161 | }; |
| 162 | } |
| 163 | } |
| 164 | |
| 165 | // Never sleep past the code's expiry, even after slow_down backoff. |
| 166 | let remaining = deadline.saturating_duration_since(Instant::now()); |
| 167 | if remaining.is_zero() { |
| 168 | break; |
| 169 | } |
| 170 | sleep(interval.min(remaining)); |
| 171 | } |
| 172 | |
| 173 | Err(self.timed_out(saw_slow_down)) |
| 174 | } |
| 175 | |
| 176 | fn timed_out(&self, saw_slow_down: bool) -> anyhow::Error { |
| 177 | match (saw_slow_down, self.slow_down_timeout_message.as_deref()) { |
| 178 | (true, Some(message)) => anyhow::anyhow!("{message}"), |
| 179 | _ => anyhow::anyhow!("{}", self.timeout_message), |
| 180 | } |
| 181 | } |
| 182 | } |
| 183 | |
| 184 | /// Reject a device-code verification URI that must not be handed to a browser |
| 185 | /// opener. |
| 186 | /// |
| 187 | /// Ported from pi's `validateVerificationUri` |
| 188 | /// (`packages/ai/src/auth/oauth/xai.ts`, MIT, Copyright (c) 2025 Mario |
| 189 | /// Zechner): the URI comes straight off the wire and is passed to the platform |
| 190 | /// "open this" call, so a malicious or compromised response could otherwise |
| 191 | /// launch `file:`, a custom app scheme, or a helper with attacker-chosen |
| 192 | /// arguments. pi requires `https:`; Codewhale additionally allows `http:` on a |
| 193 | /// loopback host, which is what self-hosted issuers and the device-code tests |
| 194 | /// use — matching the loopback allowance the account login already makes. |
| 195 | /// |
| 196 | /// Embedded credentials are rejected in every case. |
| 197 | pub fn validate_browser_verification_uri(raw: &str, context: &str) -> Result<String> { |
| 198 | let trimmed = raw.trim(); |
| 199 | let Ok(url) = url_scheme_and_host(trimmed) else { |
| 200 | bail!("{context} returned an unusable verification URI"); |
| 201 | }; |
| 202 | let (scheme, host, has_credentials) = url; |
| 203 | if has_credentials { |
| 204 | bail!("{context} returned a verification URI with embedded credentials"); |
| 205 | } |
| 206 | let allowed = scheme == "https" || (scheme == "http" && is_loopback_host(&host)); |
| 207 | if !allowed { |
| 208 | bail!("{context} returned an untrusted verification URI"); |
| 209 | } |
| 210 | Ok(trimmed.to_string()) |
| 211 | } |
| 212 | |
| 213 | /// Minimal scheme/host/credential split, so this module stays free of a URL |
| 214 | /// dependency (`codewhale-config` deliberately has no `reqwest`/`url`). |
| 215 | pub(crate) fn url_scheme_and_host(raw: &str) -> Result<(String, String, bool), ()> { |
| 216 | let (scheme, rest) = raw.split_once("://").ok_or(())?; |
| 217 | if scheme.is_empty() |
| 218 | || !scheme |
| 219 | .bytes() |
| 220 | .all(|b| b.is_ascii_alphanumeric() || b == b'+' || b == b'-' || b == b'.') |
| 221 | { |
| 222 | return Err(()); |
| 223 | } |
| 224 | let authority = rest |
| 225 | .split(['/', '?', '#']) |
| 226 | .next() |
| 227 | .filter(|authority| !authority.is_empty()) |
| 228 | .ok_or(())?; |
| 229 | let (credentials, hostport) = match authority.rsplit_once('@') { |
| 230 | Some((credentials, hostport)) => (!credentials.is_empty(), hostport), |
| 231 | None => (false, authority), |
| 232 | }; |
| 233 | let host = match hostport.strip_prefix('[') { |
| 234 | // IPv6 literal: [::1]:8080 |
| 235 | Some(rest) => rest.split_once(']').ok_or(())?.0.to_string(), |
| 236 | None => hostport.split(':').next().ok_or(())?.to_string(), |
| 237 | }; |
| 238 | if host.is_empty() { |
| 239 | return Err(()); |
| 240 | } |
| 241 | Ok(( |
| 242 | scheme.to_ascii_lowercase(), |
| 243 | host.to_ascii_lowercase(), |
| 244 | credentials, |
| 245 | )) |
| 246 | } |
| 247 | |
| 248 | pub(crate) fn is_loopback_host(host: &str) -> bool { |
| 249 | if host == "localhost" || host == "::1" { |
| 250 | return true; |
| 251 | } |
| 252 | host.parse::<std::net::IpAddr>() |
| 253 | .is_ok_and(|address| address.is_loopback()) |
| 254 | } |
| 255 | |
| 256 | #[cfg(test)] |
| 257 | mod tests { |
| 258 | use super::*; |
| 259 | use std::cell::RefCell; |
| 260 | |
| 261 | fn recording_sleep(log: &RefCell<Vec<Duration>>) -> impl FnMut(Duration) + '_ { |
| 262 | move |duration| log.borrow_mut().push(duration) |
| 263 | } |
| 264 | |
| 265 | #[test] |
| 266 | fn completes_on_first_poll_without_waiting() { |
| 267 | let slept = RefCell::new(Vec::new()); |
| 268 | let value = DeviceCodePoll::new(Duration::from_secs(60), "timed out") |
| 269 | .run(recording_sleep(&slept), || { |
| 270 | Ok(DevicePollOutcome::Complete("token")) |
| 271 | }) |
| 272 | .expect("first poll completes"); |
| 273 | assert_eq!(value, "token"); |
| 274 | assert!(slept.borrow().is_empty(), "no sleep before the first poll"); |
| 275 | } |
| 276 | |
| 277 | /// `expires_in` and `interval` are untrusted: a hostile maximum must not |
| 278 | /// panic the deadline (`Instant + Duration`) or the `slow_down` step. |
| 279 | #[test] |
| 280 | fn untrusted_maximum_lifetime_and_interval_do_not_panic() { |
| 281 | let slept = RefCell::new(Vec::new()); |
| 282 | let mut polls = 0; |
| 283 | let value = DeviceCodePoll::new(Duration::MAX, "timed out") |
| 284 | .interval_seconds(Some(u64::MAX)) |
| 285 | .run(recording_sleep(&slept), || { |
| 286 | polls += 1; |
| 287 | Ok(if polls == 1 { |
| 288 | DevicePollOutcome::SlowDown { |
| 289 | interval_seconds: None, |
| 290 | } |
| 291 | } else { |
| 292 | DevicePollOutcome::Complete("token") |
| 293 | }) |
| 294 | }) |
| 295 | .expect("a clamped run still completes"); |
| 296 | assert_eq!(value, "token"); |
| 297 | assert!(slept.borrow().iter().all(|d| *d <= MAX_LIFETIME)); |
| 298 | } |
| 299 | |
| 300 | #[test] |
| 301 | fn waits_one_interval_before_the_first_poll_when_asked() { |
| 302 | let slept = RefCell::new(Vec::new()); |
| 303 | DeviceCodePoll::new(Duration::from_secs(60), "timed out") |
| 304 | .interval_seconds(Some(3)) |
| 305 | .wait_before_first_poll(true) |
| 306 | .run(recording_sleep(&slept), || { |
| 307 | Ok(DevicePollOutcome::Complete(())) |
| 308 | }) |
| 309 | .expect("completes after the initial wait"); |
| 310 | assert_eq!(slept.borrow().as_slice(), [Duration::from_secs(3)]); |
| 311 | } |
| 312 | |
| 313 | #[test] |
| 314 | fn omitted_interval_uses_the_rfc_default_of_five_seconds() { |
| 315 | let slept = RefCell::new(Vec::new()); |
| 316 | let mut polls = 0; |
| 317 | DeviceCodePoll::new(Duration::from_secs(600), "timed out") |
| 318 | .interval_seconds(None) |
| 319 | .run(recording_sleep(&slept), || { |
| 320 | polls += 1; |
| 321 | if polls == 1 { |
| 322 | Ok(DevicePollOutcome::Pending) |
| 323 | } else { |
| 324 | Ok(DevicePollOutcome::Complete(())) |
| 325 | } |
| 326 | }) |
| 327 | .expect("completes"); |
| 328 | assert_eq!(slept.borrow().as_slice(), [Duration::from_secs(5)]); |
| 329 | } |
| 330 | |
| 331 | #[test] |
| 332 | fn slow_down_without_an_interval_adds_five_seconds() { |
| 333 | let slept = RefCell::new(Vec::new()); |
| 334 | let mut polls = 0; |
| 335 | DeviceCodePoll::new(Duration::from_secs(600), "timed out") |
| 336 | .interval_seconds(Some(2)) |
| 337 | .run(recording_sleep(&slept), || { |
| 338 | polls += 1; |
| 339 | match polls { |
| 340 | 1 => Ok(DevicePollOutcome::Pending), |
| 341 | 2 => Ok(DevicePollOutcome::SlowDown { |
| 342 | interval_seconds: None, |
| 343 | }), |
| 344 | _ => Ok(DevicePollOutcome::Complete(())), |
| 345 | } |
| 346 | }) |
| 347 | .expect("completes"); |
| 348 | assert_eq!( |
| 349 | slept.borrow().as_slice(), |
| 350 | [Duration::from_secs(2), Duration::from_secs(7)] |
| 351 | ); |
| 352 | } |
| 353 | |
| 354 | #[test] |
| 355 | fn slow_down_prefers_a_server_supplied_interval() { |
| 356 | // The clock-drift fix: the server's new minimum wins over the |
| 357 | // client-tracked interval, in both directions. |
| 358 | let slept = RefCell::new(Vec::new()); |
| 359 | let mut polls = 0; |
| 360 | DeviceCodePoll::new(Duration::from_secs(600), "timed out") |
| 361 | .interval_seconds(Some(2)) |
| 362 | .run(recording_sleep(&slept), || { |
| 363 | polls += 1; |
| 364 | match polls { |
| 365 | 1 => Ok(DevicePollOutcome::SlowDown { |
| 366 | interval_seconds: Some(30), |
| 367 | }), |
| 368 | _ => Ok(DevicePollOutcome::Complete(())), |
| 369 | } |
| 370 | }) |
| 371 | .expect("completes"); |
| 372 | assert_eq!(slept.borrow().as_slice(), [Duration::from_secs(30)]); |
| 373 | } |
| 374 | |
| 375 | #[test] |
| 376 | fn interval_never_drops_below_one_second_or_exceeds_the_cap() { |
| 377 | let slept = RefCell::new(Vec::new()); |
| 378 | let mut polls = 0; |
| 379 | DeviceCodePoll::new(Duration::from_secs(600), "timed out") |
| 380 | .interval_seconds(Some(0)) |
| 381 | .max_interval_seconds(10) |
| 382 | .run(recording_sleep(&slept), || { |
| 383 | polls += 1; |
| 384 | match polls { |
| 385 | 1 => Ok(DevicePollOutcome::SlowDown { |
| 386 | interval_seconds: Some(99), |
| 387 | }), |
| 388 | _ => Ok(DevicePollOutcome::Complete(())), |
| 389 | } |
| 390 | }) |
| 391 | .expect("completes"); |
| 392 | // interval 0 falls back to the RFC default (5s), capped at 10s. |
| 393 | assert_eq!(slept.borrow().as_slice(), [Duration::from_secs(10)]); |
| 394 | } |
| 395 | |
| 396 | #[test] |
| 397 | fn never_sleeps_past_the_deadline() { |
| 398 | let slept = RefCell::new(Vec::new()); |
| 399 | let error = DeviceCodePoll::new(Duration::from_millis(30), "timed out") |
| 400 | .interval_seconds(Some(600)) |
| 401 | .run( |
| 402 | |duration| { |
| 403 | slept.borrow_mut().push(duration); |
| 404 | std::thread::sleep(duration); |
| 405 | }, |
| 406 | || Ok(DevicePollOutcome::<()>::Pending), |
| 407 | ) |
| 408 | .expect_err("deadline stops the loop"); |
| 409 | assert_eq!(error.to_string(), "timed out"); |
| 410 | for duration in slept.borrow().iter() { |
| 411 | assert!( |
| 412 | *duration <= Duration::from_millis(30), |
| 413 | "slept {duration:?} past a 30ms deadline" |
| 414 | ); |
| 415 | } |
| 416 | } |
| 417 | |
| 418 | #[test] |
| 419 | fn a_terminal_poll_error_stops_immediately() { |
| 420 | let slept = RefCell::new(Vec::new()); |
| 421 | let error = DeviceCodePoll::new(Duration::from_secs(600), "timed out") |
| 422 | .run(recording_sleep(&slept), || { |
| 423 | Err::<DevicePollOutcome<()>, _>(anyhow::anyhow!("access_denied")) |
| 424 | }) |
| 425 | .expect_err("terminal errors propagate"); |
| 426 | assert_eq!(error.to_string(), "access_denied"); |
| 427 | assert!(slept.borrow().is_empty()); |
| 428 | } |
| 429 | |
| 430 | #[test] |
| 431 | fn timing_out_after_slow_down_reports_the_clock_drift_message() { |
| 432 | let error = DeviceCodePoll::new(Duration::from_millis(5), "plain timeout") |
| 433 | .interval_seconds(Some(1)) |
| 434 | .slow_down_timeout_message("clock drift timeout") |
| 435 | .run(std::thread::sleep, || { |
| 436 | Ok(DevicePollOutcome::<()>::SlowDown { |
| 437 | interval_seconds: None, |
| 438 | }) |
| 439 | }) |
| 440 | .expect_err("deadline stops the loop"); |
| 441 | assert_eq!(error.to_string(), "clock drift timeout"); |
| 442 | } |
| 443 | |
| 444 | #[test] |
| 445 | fn timing_out_without_slow_down_reports_the_plain_message() { |
| 446 | let error = DeviceCodePoll::new(Duration::from_millis(5), "plain timeout") |
| 447 | .interval_seconds(Some(1)) |
| 448 | .slow_down_timeout_message("clock drift timeout") |
| 449 | .run(std::thread::sleep, || Ok(DevicePollOutcome::<()>::Pending)) |
| 450 | .expect_err("deadline stops the loop"); |
| 451 | assert_eq!(error.to_string(), "plain timeout"); |
| 452 | } |
| 453 | |
| 454 | #[test] |
| 455 | fn verification_uri_must_be_https_or_loopback_http() { |
| 456 | assert_eq!( |
| 457 | validate_browser_verification_uri("https://accounts.x.ai/device", "xAI").unwrap(), |
| 458 | "https://accounts.x.ai/device" |
| 459 | ); |
| 460 | assert!(validate_browser_verification_uri("http://127.0.0.1:8080/verify", "xAI").is_ok()); |
| 461 | assert!(validate_browser_verification_uri("http://localhost/verify", "xAI").is_ok()); |
| 462 | assert!(validate_browser_verification_uri("http://[::1]:9/verify", "xAI").is_ok()); |
| 463 | |
| 464 | for hostile in [ |
| 465 | "http://accounts.x.ai/device", |
| 466 | "file:///etc/passwd", |
| 467 | "javascript:alert(1)", |
| 468 | "vscode://attacker/run", |
| 469 | "data:text/html,<script>", |
| 470 | "https://", |
| 471 | "not a url", |
| 472 | "", |
| 473 | ] { |
| 474 | assert!( |
| 475 | validate_browser_verification_uri(hostile, "xAI").is_err(), |
| 476 | "accepted {hostile}" |
| 477 | ); |
| 478 | } |
| 479 | } |
| 480 | |
| 481 | #[test] |
| 482 | fn verification_uri_rejects_embedded_credentials() { |
| 483 | let error = |
| 484 | validate_browser_verification_uri("https://user:pass@accounts.x.ai/device", "xAI") |
| 485 | .expect_err("credentials must be rejected"); |
| 486 | assert!(error.to_string().contains("embedded credentials")); |
| 487 | } |
| 488 | } |
| 489 |