M3: Copilot device-flow, token exchange, provider
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.
This commit is contained in:
@@ -0,0 +1,209 @@
|
||||
//! 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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user