diff --git a/crates/goose/src/providers/githubcopilot.rs b/crates/goose/src/providers/githubcopilot.rs index 2ce16cf2cbb6..34211075f66a 100644 --- a/crates/goose/src/providers/githubcopilot.rs +++ b/crates/goose/src/providers/githubcopilot.rs @@ -1,5 +1,6 @@ use crate::config::paths::Paths; use crate::providers::api_client::{ApiClient, AuthMethod}; +use crate::providers::oauth_device_flow::{run_device_flow, DeviceFlowConfig, RequestEncoding}; use crate::providers::openai_compatible::{handle_status_openai_compat, stream_openai_compat}; use anyhow::{anyhow, Context, Result}; use async_trait::async_trait; @@ -92,13 +93,6 @@ impl GithubCopilotUrls { } } -#[derive(Debug, Deserialize)] -struct DeviceCodeInfo { - device_code: String, - user_code: String, - verification_uri: String, -} - #[derive(Debug, Serialize, Deserialize, Clone)] struct CopilotTokenEndpoints { api: String, @@ -344,105 +338,16 @@ impl GithubCopilotProvider { } async fn login(&self) -> Result { - let device_code_info = self.get_device_code().await?; - - if let Ok(mut clipboard) = arboard::Clipboard::new() { - if let Err(e) = clipboard.set_text(&device_code_info.user_code) { - tracing::warn!("Failed to copy verification code to clipboard: {}", e); - } - } - - if let Err(e) = webbrowser::open(&device_code_info.verification_uri) { - tracing::warn!("Failed to open browser: {}", e); - } - - println!( - "Please visit {} and enter code {}", - device_code_info.verification_uri, device_code_info.user_code - ); - - self.poll_for_access_token(&device_code_info.device_code) - .await - } - - async fn get_device_code(&self) -> Result { - #[derive(Serialize)] - struct DeviceCodeRequest { - client_id: String, - scope: String, - } - self.client - .post(&self.urls.device_code_url) - .headers(self.get_github_headers()) - .json(&DeviceCodeRequest { - client_id: self.client_id.clone(), - scope: "read:user".to_string(), - }) - .send() - .await - .context("failed to send request to get device code")? - .error_for_status() - .context("failed to get device code")? - .json::() - .await - .context("failed to parse device code response") - } - - async fn poll_for_access_token(&self, device_code: &str) -> Result { - #[derive(Serialize)] - struct AccessTokenRequest { - client_id: String, - device_code: String, - grant_type: String, - } - #[derive(Debug, Deserialize)] - struct AccessTokenResponse { - access_token: Option, - error: Option, - #[serde(flatten)] - _extra: HashMap, - } - - const MAX_ATTEMPTS: i32 = 36; - for attempt in 0..MAX_ATTEMPTS { - let resp = self - .client - .post(&self.urls.access_token_url) - .headers(self.get_github_headers()) - .json(&AccessTokenRequest { - client_id: self.client_id.clone(), - device_code: device_code.to_string(), - grant_type: "urn:ietf:params:oauth:grant-type:device_code".to_string(), - }) - .send() - .await - .context("failed to make request while polling for access token")? - .error_for_status() - .context("error polling for access token")? - .json::() - .await - .context("failed to parse response while polling for access token")?; - if resp.access_token.is_some() { - tracing::trace!("successful authorization: {:#?}", resp,); - } - if let Some(access_token) = resp.access_token { - return Ok(access_token); - } else if resp - .error - .as_ref() - .is_some_and(|err| err == "authorization_pending") - { - tracing::debug!( - "authorization pending (attempt {}/{})", - attempt + 1, - MAX_ATTEMPTS - ); - } else { - tracing::debug!("unexpected response: {:#?}", resp); - } - tokio::time::sleep(tokio::time::Duration::from_secs(5)).await; - } - Err(anyhow!("failed to get access token")) + let cfg = DeviceFlowConfig { + device_auth_url: Some(&self.urls.device_code_url), + token_url: &self.urls.access_token_url, + client_id: &self.client_id, + scopes: Some("read:user"), + extra_headers: self.get_github_headers(), + encoding: RequestEncoding::Json, + }; + let tokens = run_device_flow(&self.client, &cfg).await?; + Ok(tokens.access_token) } fn get_github_headers(&self) -> http::HeaderMap { diff --git a/crates/goose/src/providers/kimicode.rs b/crates/goose/src/providers/kimicode.rs index 88065d888c53..2187d739d122 100644 --- a/crates/goose/src/providers/kimicode.rs +++ b/crates/goose/src/providers/kimicode.rs @@ -1,7 +1,7 @@ use crate::config::paths::Paths; use crate::config::Config; use crate::session_context::SESSION_ID_HEADER; -use anyhow::{anyhow, Context, Result}; +use anyhow::Result; use async_stream::try_stream; use async_trait::async_trait; use chrono::{DateTime, Duration, Utc}; @@ -19,6 +19,9 @@ use uuid::Uuid; use super::base::{ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata}; use super::errors::ProviderError; use super::formats::anthropic::{create_request, response_to_streaming_message}; +use super::oauth_device_flow::{ + refresh_device_flow_token, run_device_flow, DeviceFlowConfig, DeviceFlowTokens, RequestEncoding, +}; use super::openai_compatible::handle_status_openai_compat; use super::retry::ProviderRetry; use super::utils::RequestLog; @@ -49,16 +52,6 @@ const REFRESH_THRESHOLD_SECS: i64 = 300; /// Fallback access-token lifetime when the server omits `expires_in`. const DEFAULT_TOKEN_LIFETIME_SECS: i64 = 3600; -/// Fallback device-code window when the server omits `expires_in` -/// from `device_authorization`. -const DEFAULT_DEVICE_CODE_LIFETIME_SECS: u64 = 300; - -/// Fallback poll interval when the server omits `interval`. -const DEFAULT_POLL_INTERVAL_SECS: u64 = 5; - -/// Extra seconds added to the poll interval after an RFC 8628 `slow_down`. -const SLOW_DOWN_BACKOFF_SECS: u64 = 5; - /// Marker key written to the user config when OAuth completes successfully. /// `check_provider_configured` (server) keys off this when an OAuth-flow /// provider has no required secret env var. @@ -73,6 +66,24 @@ struct KimiToken { expires_at: DateTime, } +/// Normalize helper output into the on-disk `KimiToken` shape. When the helper +/// returns `None` for `refresh_token` or `expires_at`, fall back to the prior +/// refresh token (per RFC 6749 §6) and a default lifetime. +fn tokens_to_kimi(tokens: DeviceFlowTokens, prior_refresh: Option<&str>) -> KimiToken { + let refresh_token = tokens + .refresh_token + .or_else(|| prior_refresh.map(str::to_string)) + .unwrap_or_default(); + let expires_at = tokens + .expires_at + .unwrap_or_else(|| Utc::now() + Duration::seconds(DEFAULT_TOKEN_LIFETIME_SECS)); + KimiToken { + access_token: tokens.access_token, + refresh_token, + expires_at, + } +} + #[derive(Debug)] struct TokenCache { path: std::path::PathBuf, @@ -261,201 +272,34 @@ impl KimiCodeProvider { } async fn device_flow_login(&self) -> Result { - #[derive(Serialize)] - struct DeviceAuthReq<'a> { - client_id: &'a str, - } - #[derive(Deserialize)] - struct DeviceAuthResp { - device_code: String, - user_code: String, - verification_uri_complete: Option, - verification_uri: String, - interval: Option, - expires_in: Option, - } - - let resp: DeviceAuthResp = self - .client - .post(format!("{}/api/oauth/device_authorization", self.auth_host)) - .headers(self.kimi_headers()) - .form(&DeviceAuthReq { - client_id: KIMI_CODE_CLIENT_ID, - }) - .send() - .await - .context("failed to request device authorization")? - .error_for_status() - .context("device authorization request failed")? - .json() - .await - .context("failed to parse device authorization response")?; - - let verify_url = resp - .verification_uri_complete - .as_deref() - .unwrap_or(&resp.verification_uri); - let interval = resp.interval.unwrap_or(DEFAULT_POLL_INTERVAL_SECS); - - if let Ok(mut clipboard) = arboard::Clipboard::new() { - let _ = clipboard.set_text(&resp.user_code); - } - if let Err(e) = webbrowser::open(verify_url) { - tracing::warn!("Failed to open browser: {}", e); - } - - // stderr so CLI workflows parsing stdout aren't interfered with. - eprintln!( - "Please visit {} and enter code {}", - verify_url, resp.user_code - ); - - let expires_in = resp.expires_in.unwrap_or(DEFAULT_DEVICE_CODE_LIFETIME_SECS); - self.poll_for_token(&resp.device_code, interval, expires_in) - .await - } - - async fn poll_for_token( - &self, - device_code: &str, - interval_secs: u64, - expires_in_secs: u64, - ) -> Result { - #[derive(Serialize)] - struct PollReq<'a> { - client_id: &'a str, - device_code: &'a str, - grant_type: &'static str, - } - #[derive(Deserialize, Debug)] - struct PollResp { - access_token: Option, - refresh_token: Option, - expires_in: Option, - error: Option, - } - - let deadline = - tokio::time::Instant::now() + tokio::time::Duration::from_secs(expires_in_secs); - let mut effective_interval = interval_secs; - loop { - if tokio::time::Instant::now() >= deadline { - return Err(anyhow!("timed out waiting for user authorization")); - } - tokio::time::sleep(tokio::time::Duration::from_secs(effective_interval)).await; - - let response = self - .client - .post(format!("{}/api/oauth/token", self.auth_host)) - .headers(self.kimi_headers()) - .form(&PollReq { - client_id: KIMI_CODE_CLIENT_ID, - device_code, - grant_type: "urn:ietf:params:oauth:grant-type:device_code", - }) - .send() - .await - .context("failed to poll for token")?; - - // RFC 8628 returns pending/slow_down as 4xx with a JSON error payload, - // so don't `error_for_status()` before parsing — but if the body is - // unparseable AND the status is non-2xx, surface the HTTP status. - let status = response.status(); - let bytes = response - .bytes() - .await - .context("failed to read token poll response")?; - let resp: PollResp = match serde_json::from_slice(&bytes) { - Ok(p) => p, - Err(e) => { - if !status.is_success() { - return Err(anyhow!( - "token poll HTTP {}: {}", - status, - String::from_utf8_lossy(&bytes) - )); - } - return Err( - anyhow::Error::new(e).context("failed to parse token poll response") - ); - } - }; - - if let Some(access_token) = resp.access_token { - // RFC 6749: refresh_token is optional in token responses. - // Kimi currently returns one, but be defensive for servers/ - // versions that do not. - let refresh_token = resp.refresh_token.unwrap_or_default(); - let expires_in = resp.expires_in.unwrap_or(DEFAULT_TOKEN_LIFETIME_SECS); - return Ok(KimiToken { - access_token, - refresh_token, - expires_at: Utc::now() + Duration::seconds(expires_in), - }); - } - - match resp.error.as_deref() { - Some("authorization_pending") => { - tracing::debug!("authorization pending, continuing to poll"); - } - // RFC 8628: client MUST increase polling interval by 5 seconds - Some("slow_down") => { - tracing::debug!("slow_down received, increasing poll interval"); - effective_interval += SLOW_DOWN_BACKOFF_SECS; - } - Some(err) => { - return Err(anyhow!("authorization failed: {}", err)); - } - None => { - tracing::debug!("unexpected poll response: no token and no error"); - } - } - } + let device_auth_url = format!("{}/api/oauth/device_authorization", self.auth_host); + let token_url = format!("{}/api/oauth/token", self.auth_host); + let cfg = DeviceFlowConfig { + device_auth_url: Some(&device_auth_url), + token_url: &token_url, + client_id: KIMI_CODE_CLIENT_ID, + scopes: None, + extra_headers: self.kimi_headers(), + encoding: RequestEncoding::Form, + }; + let tokens = run_device_flow(&self.client, &cfg).await?; + Ok(tokens_to_kimi(tokens, None)) } async fn do_refresh_token(&self, refresh_token: &str) -> Result { - #[derive(Serialize)] - struct RefreshReq<'a> { - client_id: &'a str, - grant_type: &'static str, - refresh_token: &'a str, - } - #[derive(Deserialize)] - struct RefreshResp { - access_token: String, - refresh_token: Option, - expires_in: Option, - } - - let resp: RefreshResp = self - .client - .post(format!("{}/api/oauth/token", self.auth_host)) - .headers(self.kimi_headers()) - .form(&RefreshReq { - client_id: KIMI_CODE_CLIENT_ID, - grant_type: "refresh_token", - refresh_token, - }) - .send() - .await - .context("failed to refresh token")? - .error_for_status() - .context("token refresh failed")? - .json() - .await - .context("failed to parse token refresh response")?; - + let token_url = format!("{}/api/oauth/token", self.auth_host); + let cfg = DeviceFlowConfig { + device_auth_url: None, + token_url: &token_url, + client_id: KIMI_CODE_CLIENT_ID, + scopes: None, + extra_headers: self.kimi_headers(), + encoding: RequestEncoding::Form, + }; + let tokens = refresh_device_flow_token(&self.client, &cfg, refresh_token).await?; // RFC 6749 §6: the server MAY omit `refresh_token` from a refresh // response, in which case the client should keep reusing the prior one. - let next_refresh_token = resp - .refresh_token - .unwrap_or_else(|| refresh_token.to_string()); - let expires_in = resp.expires_in.unwrap_or(DEFAULT_TOKEN_LIFETIME_SECS); - Ok(KimiToken { - access_token: resp.access_token, - refresh_token: next_refresh_token, - expires_at: Utc::now() + Duration::seconds(expires_in), - }) + Ok(tokens_to_kimi(tokens, Some(refresh_token))) } // ── HTTP ───────────────────────────────────────────────────────────────── @@ -808,55 +652,10 @@ mod tests { assert_eq!(usable.refresh_token, "new_refresh"); } - #[tokio::test] - async fn poll_for_token_handles_authorization_pending_then_success() { - let server = MockServer::start().await; - - // First call: authorization_pending (returned as 400 per RFC 8628). - Mock::given(method("POST")) - .and(path("/api/oauth/token")) - .respond_with(ResponseTemplate::new(400).set_body_json(json!({ - "error": "authorization_pending", - }))) - .up_to_n_times(1) - .mount(&server) - .await; - - // Subsequent call: token issued. - Mock::given(method("POST")) - .and(path("/api/oauth/token")) - .respond_with(ResponseTemplate::new(200).set_body_json(json!({ - "access_token": "the_token", - "refresh_token": "the_refresh", - "expires_in": 1800, - }))) - .mount(&server) - .await; - - let provider = test_provider(&server.uri(), "abc"); - let token = provider.poll_for_token("device-abc", 0, 30).await.unwrap(); - assert_eq!(token.access_token, "the_token"); - assert_eq!(token.refresh_token, "the_refresh"); - } - - #[tokio::test] - async fn poll_for_token_accepts_response_without_refresh_token() { - // RFC 6749: refresh_token is optional in token responses. - let server = MockServer::start().await; - Mock::given(method("POST")) - .and(path("/api/oauth/token")) - .respond_with(ResponseTemplate::new(200).set_body_json(json!({ - "access_token": "access_only", - "expires_in": 1800, - }))) - .mount(&server) - .await; - - let provider = test_provider(&server.uri(), "abc"); - let token = provider.poll_for_token("device-abc", 0, 5).await.unwrap(); - assert_eq!(token.access_token, "access_only"); - assert_eq!(token.refresh_token, ""); - } + // NOTE: RFC 8628 polling behavior (authorization_pending, slow_down, missing + // refresh_token, HTTP errors during polling) is covered by + // `providers::oauth_device_flow` tests. Tests here focus on Kimi-specific + // integration — token cache, refresh-fallback when server omits refresh_token. #[tokio::test] async fn use_or_refresh_preserves_refresh_token_when_server_omits_it() { @@ -885,29 +684,6 @@ mod tests { assert_eq!(usable.refresh_token, "original_refresh"); } - #[tokio::test] - async fn poll_for_token_surfaces_http_error_on_unparseable_body() { - let server = MockServer::start().await; - Mock::given(method("POST")) - .and(path("/api/oauth/token")) - .respond_with(ResponseTemplate::new(502).set_body_string("Bad Gateway")) - .mount(&server) - .await; - - let provider = test_provider(&server.uri(), "abc"); - let err = provider - .poll_for_token("device-abc", 0, 5) - .await - .unwrap_err(); - let msg = format!("{:#}", err); - assert!(msg.contains("502"), "expected status in error: {}", msg); - assert!( - msg.contains("Bad Gateway"), - "expected body in error: {}", - msg - ); - } - // ── fetch_supported_models ──────────────────────────────────────────────── async fn seed_fresh_token(provider: &KimiCodeProvider) { diff --git a/crates/goose/src/providers/mod.rs b/crates/goose/src/providers/mod.rs index 616e7c49ff7d..2712a12ff028 100644 --- a/crates/goose/src/providers/mod.rs +++ b/crates/goose/src/providers/mod.rs @@ -34,6 +34,7 @@ pub mod litellm; pub mod local_inference; pub mod nanogpt; pub mod oauth; +pub mod oauth_device_flow; pub mod ollama; pub mod openai; pub mod openai_compatible; diff --git a/crates/goose/src/providers/oauth_device_flow.rs b/crates/goose/src/providers/oauth_device_flow.rs new file mode 100644 index 000000000000..d257d67b2df8 --- /dev/null +++ b/crates/goose/src/providers/oauth_device_flow.rs @@ -0,0 +1,559 @@ +//! Shared OAuth 2.0 Device Authorization Grant (RFC 8628) helper. +//! +//! Used by providers that authenticate via device-code flow (kimicode, +//! githubcopilot). Handles the authorization request, user-interaction UI, +//! polling loop with RFC 8628 `authorization_pending` / `slow_down` semantics, +//! and optional `refresh_token` grant (RFC 6749 §6). + +use anyhow::{anyhow, Context, Result}; +use chrono::{DateTime, Duration, Utc}; +use reqwest::header::HeaderMap; +use reqwest::Client; +use serde::{Deserialize, Serialize}; + +/// Fallback poll interval when the server omits `interval` (RFC 8628 §3.2). +const DEFAULT_POLL_INTERVAL_SECS: u64 = 5; + +/// Fallback device-code window when the server omits `expires_in` (RFC 8628 §3.2). +const DEFAULT_DEVICE_CODE_LIFETIME_SECS: u64 = 300; + +/// Extra seconds added to the poll interval after an RFC 8628 `slow_down`. +const SLOW_DOWN_BACKOFF_SECS: u64 = 5; + +/// How a provider expects the device-authorization and token request bodies to +/// be encoded. RFC 8628 §3.1 specifies `application/x-www-form-urlencoded`, but +/// GitHub accepts JSON when `Accept: application/json` is set. +#[derive(Debug, Clone, Copy)] +pub enum RequestEncoding { + Form, + Json, +} + +/// Connection details for a provider's device flow. +#[derive(Debug, Clone)] +pub struct DeviceFlowConfig<'a> { + /// `device_authorization_endpoint` (RFC 8628 §3.1). + /// `None` when only the refresh grant is needed. + pub device_auth_url: Option<&'a str>, + /// `token_endpoint` used for both device-code polling and refresh grants. + pub token_url: &'a str, + /// Public OAuth client identifier. + pub client_id: &'a str, + /// Space-separated scope string, or `None` to omit the parameter. + pub scopes: Option<&'a str>, + /// Provider-specific headers (user-agent, platform markers, `Accept`, etc.). + pub extra_headers: HeaderMap, + /// Body encoding for device-auth, polling, and refresh requests. + pub encoding: RequestEncoding, +} + +/// Fields returned by `/device_authorization` (RFC 8628 §3.2). +#[derive(Debug, Clone, Deserialize)] +pub struct DeviceCodeResponse { + pub device_code: String, + pub user_code: String, + pub verification_uri: String, + /// Pre-populated URI with user_code embedded, used when the provider + /// supports it (e.g. Kimi). Fall back to `verification_uri` otherwise. + pub verification_uri_complete: Option, + pub interval: Option, + pub expires_in: Option, +} + +impl DeviceCodeResponse { + /// URI the user should visit. Prefers the `_complete` form when present. + pub fn verification_url(&self) -> &str { + self.verification_uri_complete + .as_deref() + .unwrap_or(&self.verification_uri) + } +} + +/// Access + optional refresh credentials from a device-code exchange. +#[derive(Debug, Clone)] +pub struct DeviceFlowTokens { + pub access_token: String, + /// Some providers (GitHub Copilot) do not issue a refresh token. + pub refresh_token: Option, + /// Derived from `expires_in` on the token response. `None` when the server + /// omits it (RFC 6749 §5.1 permits that). + pub expires_at: Option>, +} + +// ── Public entry points ────────────────────────────────────────────────────── + +/// Request a device code from the authorization server. +pub async fn request_device_code( + client: &Client, + cfg: &DeviceFlowConfig<'_>, +) -> Result { + #[derive(Serialize)] + struct DeviceAuthReq<'a> { + client_id: &'a str, + #[serde(skip_serializing_if = "Option::is_none")] + scope: Option<&'a str>, + } + + let body = DeviceAuthReq { + client_id: cfg.client_id, + scope: cfg.scopes, + }; + + let url = cfg + .device_auth_url + .ok_or_else(|| anyhow!("device_auth_url is required for device code request"))?; + send_request(client, cfg, url, &body) + .await + .context("failed to request device authorization")? + .error_for_status() + .context("device authorization request failed")? + .json::() + .await + .context("failed to parse device authorization response") +} + +/// Poll the token endpoint until the user authorizes (or the device code expires). +/// Implements RFC 8628 §3.5 — handles `authorization_pending` and `slow_down`. +pub async fn poll_for_tokens( + client: &Client, + cfg: &DeviceFlowConfig<'_>, + device_code: &str, + interval_secs: u64, + expires_in_secs: u64, +) -> Result { + #[derive(Serialize)] + struct PollReq<'a> { + client_id: &'a str, + device_code: &'a str, + grant_type: &'static str, + } + + let req = PollReq { + client_id: cfg.client_id, + device_code, + grant_type: "urn:ietf:params:oauth:grant-type:device_code", + }; + + let deadline = tokio::time::Instant::now() + tokio::time::Duration::from_secs(expires_in_secs); + let mut effective_interval = interval_secs; + + loop { + if tokio::time::Instant::now() >= deadline { + return Err(anyhow!("timed out waiting for user authorization")); + } + tokio::time::sleep(tokio::time::Duration::from_secs(effective_interval)).await; + + let response = send_request(client, cfg, cfg.token_url, &req) + .await + .context("failed to poll for token")?; + + // RFC 8628 §3.5 returns pending/slow_down as 4xx with a JSON error + // payload, so don't `error_for_status()` before parsing. If the body + // is unparseable AND the status is non-2xx, surface the HTTP status. + match parse_token_response(response).await? { + TokenPollOutcome::Issued(tokens) => return Ok(tokens), + TokenPollOutcome::Pending => { + tracing::debug!("authorization pending, continuing to poll"); + } + TokenPollOutcome::SlowDown => { + tracing::debug!("slow_down received, increasing poll interval"); + effective_interval += SLOW_DOWN_BACKOFF_SECS; + } + TokenPollOutcome::Failed(err) => { + return Err(anyhow!("authorization failed: {}", err)); + } + } + } +} + +/// High-level flow: request a device code, print user-facing instructions, +/// open the browser, and poll until tokens are issued. +pub async fn run_device_flow( + client: &Client, + cfg: &DeviceFlowConfig<'_>, +) -> Result { + let device = request_device_code(client, cfg).await?; + announce_user_action(&device); + + let interval = device.interval.unwrap_or(DEFAULT_POLL_INTERVAL_SECS); + let expires_in = device + .expires_in + .unwrap_or(DEFAULT_DEVICE_CODE_LIFETIME_SECS); + + poll_for_tokens(client, cfg, &device.device_code, interval, expires_in).await +} + +/// Exchange a refresh token for a new access token (RFC 6749 §6). +pub async fn refresh_device_flow_token( + client: &Client, + cfg: &DeviceFlowConfig<'_>, + refresh_token: &str, +) -> Result { + #[derive(Serialize)] + struct RefreshReq<'a> { + client_id: &'a str, + grant_type: &'static str, + refresh_token: &'a str, + } + + let req = RefreshReq { + client_id: cfg.client_id, + grant_type: "refresh_token", + refresh_token, + }; + + let raw: TokenResponseBody = send_request(client, cfg, cfg.token_url, &req) + .await + .context("failed to refresh token")? + .error_for_status() + .context("token refresh failed")? + .json() + .await + .context("failed to parse token refresh response")?; + + let access_token = raw + .access_token + .ok_or_else(|| anyhow!("refresh response missing access_token"))?; + Ok(DeviceFlowTokens { + access_token, + refresh_token: raw.refresh_token, + expires_at: raw + .expires_in + .map(|secs| Utc::now() + Duration::seconds(secs)), + }) +} + +// ── Internals ──────────────────────────────────────────────────────────────── + +#[derive(Debug, Deserialize)] +struct TokenResponseBody { + access_token: Option, + refresh_token: Option, + expires_in: Option, + error: Option, +} + +enum TokenPollOutcome { + Issued(DeviceFlowTokens), + Pending, + SlowDown, + Failed(String), +} + +async fn parse_token_response(response: reqwest::Response) -> Result { + let status = response.status(); + let bytes = response + .bytes() + .await + .context("failed to read token poll response")?; + + let body: TokenResponseBody = match serde_json::from_slice(&bytes) { + Ok(p) => p, + Err(e) => { + if !status.is_success() { + return Err(anyhow!( + "token poll HTTP {}: {}", + status, + String::from_utf8_lossy(&bytes) + )); + } + return Err(anyhow::Error::new(e).context("failed to parse token poll response")); + } + }; + + if let Some(access_token) = body.access_token { + return Ok(TokenPollOutcome::Issued(DeviceFlowTokens { + access_token, + refresh_token: body.refresh_token, + expires_at: body + .expires_in + .map(|secs| Utc::now() + Duration::seconds(secs)), + })); + } + + Ok(match body.error.as_deref() { + Some("authorization_pending") => TokenPollOutcome::Pending, + Some("slow_down") => TokenPollOutcome::SlowDown, + Some(err) => TokenPollOutcome::Failed(err.to_string()), + None => TokenPollOutcome::Failed( + "unexpected token response: no access_token and no error code".to_string(), + ), + }) +} + +async fn send_request( + client: &Client, + cfg: &DeviceFlowConfig<'_>, + url: &str, + body: &T, +) -> reqwest::Result { + let builder = client.post(url).headers(cfg.extra_headers.clone()); + let builder = match cfg.encoding { + RequestEncoding::Form => builder.form(body), + RequestEncoding::Json => builder.json(body), + }; + builder.send().await +} + +fn announce_user_action(device: &DeviceCodeResponse) { + if let Ok(mut clipboard) = arboard::Clipboard::new() { + if let Err(e) = clipboard.set_text(&device.user_code) { + tracing::warn!("Failed to copy verification code to clipboard: {}", e); + } + } + let verify_url = device.verification_url(); + if let Err(e) = webbrowser::open(verify_url) { + tracing::warn!("Failed to open browser: {}", e); + } + // stderr keeps stdout clean for CLI workflows parsing provider output. + eprintln!( + "Please visit {} and enter code {}", + verify_url, device.user_code + ); +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + fn make_cfg<'a>(device_auth_url: Option<&'a str>, token_url: &'a str) -> DeviceFlowConfig<'a> { + DeviceFlowConfig { + device_auth_url, + token_url, + client_id: "test-client", + scopes: None, + extra_headers: HeaderMap::new(), + encoding: RequestEncoding::Form, + } + } + + #[tokio::test] + async fn poll_returns_issued_tokens_when_server_responds_immediately() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "the_token", + "refresh_token": "the_refresh", + "expires_in": 1800, + }))) + .mount(&server) + .await; + + let token_url = format!("{}/token", server.uri()); + let cfg = make_cfg(None, &token_url); + + let client = Client::new(); + let tokens = poll_for_tokens(&client, &cfg, "device-abc", 0, 30) + .await + .unwrap(); + assert_eq!(tokens.access_token, "the_token"); + assert_eq!(tokens.refresh_token.as_deref(), Some("the_refresh")); + assert!(tokens.expires_at.is_some()); + } + + #[tokio::test] + async fn poll_handles_authorization_pending_then_success() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(400).set_body_json(json!({ + "error": "authorization_pending", + }))) + .up_to_n_times(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "issued", + "expires_in": 900, + }))) + .mount(&server) + .await; + + let token_url = format!("{}/token", server.uri()); + let cfg = make_cfg(None, &token_url); + + let client = Client::new(); + let tokens = poll_for_tokens(&client, &cfg, "device-abc", 0, 30) + .await + .unwrap(); + assert_eq!(tokens.access_token, "issued"); + assert!(tokens.refresh_token.is_none()); + } + + #[tokio::test] + async fn poll_handles_slow_down_then_success() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(400).set_body_json(json!({ + "error": "slow_down", + }))) + .up_to_n_times(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "issued", + "expires_in": 900, + }))) + .mount(&server) + .await; + + let token_url = format!("{}/token", server.uri()); + let cfg = make_cfg(None, &token_url); + + let client = Client::new(); + let tokens = poll_for_tokens(&client, &cfg, "device-abc", 0, 30) + .await + .unwrap(); + assert_eq!(tokens.access_token, "issued"); + } + + #[tokio::test] + async fn poll_times_out_when_user_never_authorizes() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(400).set_body_json(json!({ + "error": "authorization_pending", + }))) + .mount(&server) + .await; + + let token_url = format!("{}/token", server.uri()); + let cfg = make_cfg(None, &token_url); + + let client = Client::new(); + let err = poll_for_tokens(&client, &cfg, "device-abc", 0, 0) + .await + .unwrap_err(); + assert!(err.to_string().contains("timed out"), "got: {}", err); + } + + #[tokio::test] + async fn poll_surfaces_http_status_on_unparseable_body() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(502).set_body_string("Bad Gateway")) + .mount(&server) + .await; + + let token_url = format!("{}/token", server.uri()); + let cfg = make_cfg(None, &token_url); + + let client = Client::new(); + let err = poll_for_tokens(&client, &cfg, "device-abc", 0, 5) + .await + .unwrap_err(); + let msg = format!("{:#}", err); + assert!(msg.contains("502"), "expected status in error: {}", msg); + assert!(msg.contains("Bad Gateway"), "expected body: {}", msg); + } + + #[tokio::test] + async fn poll_surfaces_server_error_message() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(400).set_body_json(json!({ + "error": "access_denied", + }))) + .mount(&server) + .await; + + let token_url = format!("{}/token", server.uri()); + let cfg = make_cfg(None, &token_url); + + let client = Client::new(); + let err = poll_for_tokens(&client, &cfg, "device-abc", 0, 5) + .await + .unwrap_err(); + assert!(err.to_string().contains("access_denied"), "got: {}", err); + } + + #[tokio::test] + async fn request_device_code_parses_complete_response() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/device_authorization")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "device_code": "dc", + "user_code": "UC-1", + "verification_uri": "https://example.com/activate", + "verification_uri_complete": "https://example.com/activate?user_code=UC-1", + "interval": 3, + "expires_in": 600, + }))) + .mount(&server) + .await; + + let device_url = format!("{}/device_authorization", server.uri()); + let cfg = make_cfg(Some(&device_url), ""); + + let client = Client::new(); + let resp = request_device_code(&client, &cfg).await.unwrap(); + assert_eq!(resp.device_code, "dc"); + assert_eq!(resp.user_code, "UC-1"); + assert_eq!( + resp.verification_url(), + "https://example.com/activate?user_code=UC-1" + ); + assert_eq!(resp.interval, Some(3)); + } + + #[tokio::test] + async fn refresh_token_returns_new_credentials() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "new_access", + "refresh_token": "new_refresh", + "expires_in": 3600, + }))) + .mount(&server) + .await; + + let token_url = format!("{}/token", server.uri()); + let cfg = make_cfg(None, &token_url); + + let client = Client::new(); + let tokens = refresh_device_flow_token(&client, &cfg, "old_refresh") + .await + .unwrap(); + assert_eq!(tokens.access_token, "new_access"); + assert_eq!(tokens.refresh_token.as_deref(), Some("new_refresh")); + } + + #[tokio::test] + async fn refresh_token_allows_server_to_omit_refresh_token() { + // RFC 6749 §6: server MAY omit refresh_token; caller should reuse prior. + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "new_access", + "expires_in": 3600, + }))) + .mount(&server) + .await; + + let token_url = format!("{}/token", server.uri()); + let cfg = make_cfg(None, &token_url); + + let client = Client::new(); + let tokens = refresh_device_flow_token(&client, &cfg, "old_refresh") + .await + .unwrap(); + assert_eq!(tokens.access_token, "new_access"); + assert!(tokens.refresh_token.is_none()); + } +}