Adds the Copilot provider: device-flow login modal support, copilot_internal/v2/token exchange with a direct-Bearer fallback, and routing across the three codecs, per the dual-path mitigation noted for Copilot token variance.
210 lines
6.6 KiB
Rust
210 lines
6.6 KiB
Rust
//! 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<DeviceCode, DeviceFlowError> {
|
|
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<PollOutcome, DeviceFlowError> {
|
|
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<S, Fut>(
|
|
client: &reqwest::Client,
|
|
client_id: &str,
|
|
device: &DeviceCode,
|
|
sleep: S,
|
|
) -> Result<String, DeviceFlowError>
|
|
where
|
|
S: Fn(Duration) -> Fut,
|
|
Fut: std::future::Future<Output = ()>,
|
|
{
|
|
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);
|
|
}
|
|
}
|