use super::{
ApiKeyRecord, AuthContext, AuthError, Authenticator, DeviceCodePrompt, refresh_api_key,
};
use crate::http_client::HttpClientExt;
use crate::providers::internal::auth::device::{
emit_device_code_prompt, ensure_parent_dir, read_json_record, token_expired, write_json_record,
};
use crate::providers::internal::auth::{request, send_json};
use bytes::Bytes;
use http::Method;
use serde::Deserialize;
use std::hash::{DefaultHasher, Hash, Hasher};
const GITHUB_CLIENT_ID: &str = "Iv1.b507a08c87ecfe98";
const GITHUB_DEVICE_CODE_URL: &str = "https://github.com/login/device/code";
const GITHUB_ACCESS_TOKEN_URL: &str = "https://github.com/login/oauth/access_token";
const DEVICE_CODE_POLL_SLEEP_SECONDS: u64 = 5;
const DEVICE_CODE_TIMEOUT_SECONDS: u64 = 15 * 60;
const DEVICE_CODE_SLOW_DOWN_SECONDS: u64 = 5;
#[derive(Debug, Deserialize)]
struct DeviceCodeResponse {
device_code: String,
user_code: String,
verification_uri: String,
interval: Option<u64>,
expires_in: Option<u64>,
}
#[derive(Debug, Deserialize)]
struct AccessTokenResponse {
access_token: Option<String>,
error: Option<String>,
error_description: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct AccessTokenState {
token: String,
from_cache: bool,
}
impl Authenticator {
pub(super) async fn auth_context_oauth<H>(&self, http: &H) -> Result<AuthContext, AuthError>
where
H: HttpClientExt,
{
let _refresh = self.refresh_lock.lock().await;
let record: ApiKeyRecord = read_json_record(self.api_key_file.as_deref())?;
let cached_access_token = self.read_access_token().ok().flatten();
if record.can_reuse_for_oauth(cached_access_token.as_deref()) {
return Ok(record.into_context());
}
let access_token = if let Some(token) = cached_access_token {
AccessTokenState {
token,
from_cache: true,
}
} else {
self.access_token(http).await?
};
let record = match refresh_api_key(http, &access_token.token).await {
Ok(record) => record.bind_to_bootstrap_token(&access_token.token),
Err(err) if access_token.from_cache && should_retry_with_fresh_access_token(&err) => {
self.clear_access_token()?;
let fresh_access_token = self.reauthenticate_access_token(http).await?;
refresh_api_key(http, &fresh_access_token)
.await?
.bind_to_bootstrap_token(&fresh_access_token)
}
Err(err) => return Err(err),
};
write_json_record(self.api_key_file.as_deref(), &record)?;
Ok(record.into_context())
}
pub(super) async fn auth_context_with_github_access_token<H>(
&self,
http: &H,
access_token: &str,
) -> Result<AuthContext, AuthError>
where
H: HttpClientExt,
{
let _refresh = self.refresh_lock.lock().await;
let record: ApiKeyRecord = read_json_record(self.api_key_file.as_deref())?;
if record.can_reuse_for_bootstrap_token(access_token) {
return Ok(record.into_context());
}
let record = refresh_api_key(http, access_token)
.await?
.bind_to_bootstrap_token(access_token);
write_json_record(self.api_key_file.as_deref(), &record)?;
Ok(record.into_context())
}
async fn access_token<H>(&self, http: &H) -> Result<AccessTokenState, AuthError>
where
H: HttpClientExt,
{
if let Some(token) = self.read_access_token()? {
return Ok(AccessTokenState {
token,
from_cache: true,
});
}
self.reauthenticate_access_token(http)
.await
.map(|token| AccessTokenState {
token,
from_cache: false,
})
}
async fn login_device_flow<H>(&self, http: &H) -> Result<String, AuthError>
where
H: HttpClientExt,
{
let body = url::form_urlencoded::Serializer::new(String::new())
.append_pair("client_id", GITHUB_CLIENT_ID)
.append_pair("scope", "read:user")
.finish();
let device: DeviceCodeResponse = send_json(
http,
request(Method::POST, GITHUB_DEVICE_CODE_URL)
.header(http::header::ACCEPT, "application/json")
.header(
http::header::CONTENT_TYPE,
"application/x-www-form-urlencoded",
)
.body(Bytes::from(body)),
)
.await?;
emit_device_code_prompt(
self.device_code_handler.0.as_ref(),
DeviceCodePrompt {
verification_uri: device.verification_uri.clone(),
user_code: device.user_code.clone(),
},
&format!(
"Sign in with GitHub Copilot:\n1) Visit {}\n2) Enter code: {}",
device.verification_uri, device.user_code
),
);
let deadline = std::time::Instant::now()
+ std::time::Duration::from_secs(
device.expires_in.unwrap_or(DEVICE_CODE_TIMEOUT_SECONDS),
);
let mut interval = normalize_poll_interval_seconds(device.interval);
while std::time::Instant::now() < deadline {
let body = url::form_urlencoded::Serializer::new(String::new())
.append_pair("client_id", GITHUB_CLIENT_ID)
.append_pair("device_code", &device.device_code)
.append_pair("grant_type", "urn:ietf:params:oauth:grant-type:device_code")
.finish();
let response: AccessTokenResponse = send_json(
http,
request(Method::POST, GITHUB_ACCESS_TOKEN_URL)
.header(http::header::ACCEPT, "application/json")
.header(
http::header::CONTENT_TYPE,
"application/x-www-form-urlencoded",
)
.body(Bytes::from(body)),
)
.await?;
if let Some(access_token) = response.access_token {
return Ok(access_token);
}
interval = next_poll_interval_seconds(
interval,
response.error.as_deref(),
response.error_description.as_deref(),
)?;
crate::wasm_compat::sleep(std::time::Duration::from_secs(interval)).await;
}
Err(AuthError::Message(
"Timed out waiting for GitHub Copilot device authorization".into(),
))
}
fn read_access_token(&self) -> Result<Option<String>, AuthError> {
let Some(path) = &self.access_token_file else {
return Ok(None);
};
match std::fs::read_to_string(path) {
Ok(token) => {
let token = token.trim();
if token.is_empty() {
Ok(None)
} else {
Ok(Some(token.to_owned()))
}
}
Err(err) if err.kind() == std::io::ErrorKind::NotFound => Ok(None),
Err(err) => Err(err.into()),
}
}
fn write_access_token(&self, token: &str) -> Result<(), AuthError> {
let Some(path) = &self.access_token_file else {
return Ok(());
};
ensure_parent_dir(path)?;
std::fs::write(path, token.as_bytes())?;
Ok(())
}
fn clear_access_token(&self) -> Result<(), AuthError> {
let Some(path) = &self.access_token_file else {
return Ok(());
};
match std::fs::remove_file(path) {
Ok(()) => Ok(()),
Err(err) if err.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(err) => Err(err.into()),
}
}
async fn reauthenticate_access_token<H>(&self, http: &H) -> Result<String, AuthError>
where
H: HttpClientExt,
{
if !self.allow_device_flow {
return Err(AuthError::Message(
"GitHub Copilot sign-in required. Reconnect Copilot in Settings before using this provider."
.into(),
));
}
let token = self.login_device_flow(http).await?;
self.write_access_token(&token)?;
Ok(token)
}
}
impl ApiKeyRecord {
fn can_reuse_for_oauth(&self, bootstrap_token: Option<&str>) -> bool {
if !self.has_live_api_key() {
return false;
}
bootstrap_token.is_none_or(|bootstrap_token| self.matches_bootstrap_token(bootstrap_token))
}
fn can_reuse_for_bootstrap_token(&self, bootstrap_token: &str) -> bool {
self.has_live_api_key() && self.matches_bootstrap_token(bootstrap_token)
}
fn bind_to_bootstrap_token(mut self, bootstrap_token: &str) -> Self {
self.bootstrap_token_fingerprint = Some(bootstrap_token_fingerprint(bootstrap_token));
self
}
fn has_live_api_key(&self) -> bool {
self.token
.as_ref()
.is_some_and(|token| !token.trim().is_empty())
&& !token_expired(self.expires_at, 0)
}
fn matches_bootstrap_token(&self, bootstrap_token: &str) -> bool {
self.bootstrap_token_fingerprint.as_deref()
== Some(bootstrap_token_fingerprint(bootstrap_token).as_str())
}
}
fn bootstrap_token_fingerprint(bootstrap_token: &str) -> String {
let mut hasher = DefaultHasher::new();
bootstrap_token.hash(&mut hasher);
format!("{:016x}", hasher.finish())
}
fn normalize_poll_interval_seconds(interval: Option<u64>) -> u64 {
interval.unwrap_or(DEVICE_CODE_POLL_SLEEP_SECONDS).max(1)
}
fn next_poll_interval_seconds(
current_interval: u64,
error: Option<&str>,
error_description: Option<&str>,
) -> Result<u64, AuthError> {
match error {
Some("authorization_pending") => Ok(current_interval),
Some("slow_down") => Ok(current_interval.saturating_add(DEVICE_CODE_SLOW_DOWN_SECONDS)),
Some("expired_token") => Err(AuthError::Message(
"GitHub device authorization expired before it completed".into(),
)),
Some("access_denied") => Err(AuthError::Message(
"GitHub device authorization was denied".into(),
)),
Some(other) => Err(AuthError::Message(format_oauth_error(
"GitHub device authorization failed",
other,
error_description,
))),
None => Err(AuthError::Message(
"GitHub device authorization failed: unknown error".into(),
)),
}
}
fn format_oauth_error(prefix: &str, error: &str, description: Option<&str>) -> String {
match description
.map(str::trim)
.filter(|description| !description.is_empty())
{
Some(description) => format!("{prefix}: {error} ({description})"),
None => format!("{prefix}: {error}"),
}
}
fn should_retry_with_fresh_access_token(err: &AuthError) -> bool {
match err {
AuthError::Http(err) => {
should_retry_with_fresh_access_token_status(err.non_success_status())
}
_ => false,
}
}
fn should_retry_with_fresh_access_token_status(status: Option<http::StatusCode>) -> bool {
matches!(
status,
Some(http::StatusCode::UNAUTHORIZED | http::StatusCode::FORBIDDEN)
)
}
#[cfg(test)]
mod tests;