rig-core 0.44.0

An opinionated library for building LLM powered applications.
Documentation
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;