返回 CodeWhale
device_code.rs
根目录 / crates / config / src / device_code.rs
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
489 lines RUST