//! GitHub OAuth device flow (RFC 8628) for Copilot login. //! //! ai-harness must register its own GitHub OAuth app and supply its client id (we do not //! hardcode opencode's). The client id is passed in by the caller — see [`request_device_code`]. use std::time::Duration; use serde::Deserialize; use serde_json::Value; const DEVICE_CODE_URL: &str = "https://github.com/login/device/code"; const ACCESS_TOKEN_URL: &str = "https://github.com/login/oauth/access_token"; const GRANT_TYPE: &str = "urn:ietf:params:oauth:grant-type:device_code"; /// Minimum scope needed to call the Copilot token-exchange endpoint. pub const SCOPE: &str = "read:user"; #[derive(Debug, Clone, Deserialize)] pub struct DeviceCode { pub device_code: String, pub user_code: String, pub verification_uri: String, /// Seconds between polls; GitHub requires honoring this and any `slow_down` bumps. #[serde(default = "default_interval")] pub interval: u64, #[serde(default)] pub expires_in: u64, } fn default_interval() -> u64 { 5 } /// Result of one poll of the access-token endpoint. #[derive(Debug, Clone, PartialEq, Eq)] pub enum PollOutcome { /// The user hasn't authorized yet; keep polling at the current interval. Pending, /// GitHub asked us to slow down; add 5s to the interval (RFC 8628). SlowDown, /// Authorization complete. Success { access_token: String }, /// Terminal failure (expired code, denied, unknown error) with a human-readable reason. Failed(String), } #[derive(Debug, thiserror::Error)] pub enum DeviceFlowError { #[error("network: {0}")] Network(String), #[error("unexpected response: {0}")] Unexpected(String), } /// Pure mapping of an access-token poll response body to a [`PollOutcome`], so the state /// machine is testable without a live GitHub. pub fn parse_poll_response(body: &Value) -> PollOutcome { if let Some(token) = body["access_token"].as_str() { return PollOutcome::Success { access_token: token.to_string(), }; } match body["error"].as_str() { Some("authorization_pending") => PollOutcome::Pending, Some("slow_down") => PollOutcome::SlowDown, Some(other) => { let desc = body["error_description"].as_str().unwrap_or(other); PollOutcome::Failed(desc.to_string()) } None => PollOutcome::Failed("no access_token and no error in response".to_string()), } } /// Applies a `slow_down` to a poll interval per RFC 8628 (+5s). pub fn bump_interval(interval: u64) -> u64 { interval + 5 } /// Step 1: request a device + user code for `client_id`. pub async fn request_device_code( client: &reqwest::Client, client_id: &str, ) -> Result { let resp = client .post(DEVICE_CODE_URL) .header("accept", "application/json") .json(&serde_json::json!({"client_id": client_id, "scope": SCOPE})) .send() .await .map_err(|e| DeviceFlowError::Network(e.to_string()))?; let value: Value = resp .json() .await .map_err(|e| DeviceFlowError::Network(e.to_string()))?; serde_json::from_value(value.clone()) .map_err(|_| DeviceFlowError::Unexpected(value.to_string())) } /// Step 2 (single poll): exchange the device code for an access token, once. pub async fn poll_once( client: &reqwest::Client, client_id: &str, device_code: &str, ) -> Result { let resp = client .post(ACCESS_TOKEN_URL) .header("accept", "application/json") .json(&serde_json::json!({ "client_id": client_id, "device_code": device_code, "grant_type": GRANT_TYPE, })) .send() .await .map_err(|e| DeviceFlowError::Network(e.to_string()))?; let value: Value = resp .json() .await .map_err(|e| DeviceFlowError::Network(e.to_string()))?; Ok(parse_poll_response(&value)) } /// Step 2 (full loop): polls until success or terminal failure, honoring `interval` and /// `slow_down`. `sleep` is injected so tests can drive it without real time. pub async fn poll_for_token( client: &reqwest::Client, client_id: &str, device: &DeviceCode, sleep: S, ) -> Result where S: Fn(Duration) -> Fut, Fut: std::future::Future, { let mut interval = device.interval; loop { sleep(Duration::from_secs(interval)).await; match poll_once(client, client_id, &device.device_code).await? { PollOutcome::Pending => {} PollOutcome::SlowDown => interval = bump_interval(interval), PollOutcome::Success { access_token } => return Ok(access_token), PollOutcome::Failed(reason) => return Err(DeviceFlowError::Unexpected(reason)), } } } #[cfg(test)] mod tests { use super::*; use serde_json::json; #[test] fn parses_success() { let outcome = parse_poll_response(&json!({"access_token": "gho_abc", "token_type": "bearer"})); assert_eq!( outcome, PollOutcome::Success { access_token: "gho_abc".into() } ); } #[test] fn parses_pending_and_slow_down() { assert_eq!( parse_poll_response(&json!({"error": "authorization_pending"})), PollOutcome::Pending ); assert_eq!( parse_poll_response(&json!({"error": "slow_down", "interval": 10})), PollOutcome::SlowDown ); } #[test] fn parses_terminal_errors_with_description() { match parse_poll_response( &json!({"error": "expired_token", "error_description": "code expired"}), ) { PollOutcome::Failed(msg) => assert_eq!(msg, "code expired"), other => panic!("expected Failed, got {other:?}"), } assert!(matches!( parse_poll_response(&json!({"error": "access_denied"})), PollOutcome::Failed(_) )); assert!(matches!( parse_poll_response(&json!({})), PollOutcome::Failed(_) )); } #[test] fn slow_down_adds_five_seconds() { assert_eq!(bump_interval(5), 10); assert_eq!(bump_interval(10), 15); } #[test] fn device_code_defaults_interval() { let dc: DeviceCode = serde_json::from_value(json!({ "device_code": "d", "user_code": "WXYZ-1234", "verification_uri": "https://github.com/login/device" })) .unwrap(); assert_eq!(dc.interval, 5); } }