Skip to main content

rig_core/providers/copilot/auth/
mod.rs

1use crate::http_client::HttpClientExt;
2use crate::providers::internal::auth::{request, send_json};
3use crate::wire::Secret;
4use futures::lock::Mutex;
5use http::Method;
6use serde::{Deserialize, Serialize};
7use std::fmt;
8use std::path::PathBuf;
9use std::sync::Arc;
10
11pub use crate::providers::internal::auth::{DeviceCodeHandler, DeviceCodePrompt};
12
13#[cfg(not(target_family = "wasm"))]
14mod native;
15
16const GITHUB_API_KEY_URL: &str = "https://api.github.com/copilot_internal/v2/token";
17
18/// Return `{config_dir}/github_copilot`, or `None` without a platform config directory.
19/// The directory conventionally contains `access-token` and `api-key.json`.
20pub fn default_token_dir() -> Option<PathBuf> {
21    crate::providers::internal::auth::config_dir().map(|dir| dir.join("github_copilot"))
22}
23
24/// Derive the Copilot REST base URL from a chat token's `proxy-ep=` segment.
25///
26/// The endpoint is parsed from a credential string, not from explicit caller
27/// configuration. For that reason, token-derived routing is limited to GitHub
28/// Copilot service hosts and HTTPS. Callers that need a custom non-GitHub host
29/// can still opt in explicitly with [`CopilotConfig::with_base_url`](super::CopilotConfig::with_base_url).
30pub(crate) fn base_url_from_token(token: &str) -> Option<String> {
31    let proxy_ep = token
32        .split(';')
33        .find_map(|part| part.trim().strip_prefix("proxy-ep="))?
34        .trim();
35
36    normalize_copilot_proxy_endpoint(proxy_ep)
37}
38
39fn normalize_copilot_proxy_endpoint(proxy_ep: &str) -> Option<String> {
40    if proxy_ep.is_empty() {
41        return None;
42    }
43
44    let candidate = if proxy_ep.starts_with("http://") || proxy_ep.starts_with("https://") {
45        proxy_ep.to_string()
46    } else {
47        format!("https://{proxy_ep}")
48    };
49
50    let mut url = url::Url::parse(&candidate).ok()?;
51    if url.scheme() != "https" || !url.username().is_empty() || url.password().is_some() {
52        return None;
53    }
54    if url.path() != "/" || url.query().is_some() || url.fragment().is_some() {
55        return None;
56    }
57
58    let host = url.host_str()?.to_ascii_lowercase();
59    if !is_allowed_token_derived_copilot_host(&host) {
60        return None;
61    }
62
63    let api_host = host
64        .strip_prefix("proxy.")
65        .map(|suffix| format!("api.{suffix}"))
66        .unwrap_or(host);
67    url.set_host(Some(&api_host)).ok()?;
68
69    Some(url.to_string().trim_end_matches('/').to_string())
70}
71
72fn is_allowed_token_derived_copilot_host(host: &str) -> bool {
73    host == "githubcopilot.com" || host.ends_with(".githubcopilot.com")
74}
75
76#[derive(Clone)]
77pub enum AuthSource {
78    ApiKey(String),
79    GitHubAccessToken(String),
80    OAuth,
81}
82
83impl fmt::Debug for AuthSource {
84    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
85        match self {
86            Self::ApiKey(_) => f.write_str("ApiKey(<redacted>)"),
87            Self::GitHubAccessToken(_) => f.write_str("GitHubAccessToken(<redacted>)"),
88            Self::OAuth => f.write_str("OAuth"),
89        }
90    }
91}
92
93#[derive(Clone, Debug)]
94#[cfg_attr(target_family = "wasm", allow(dead_code))]
95pub struct Authenticator {
96    source: AuthSource,
97    access_token_file: Option<PathBuf>,
98    api_key_file: Option<PathBuf>,
99    device_code_handler: DeviceCodeHandler,
100    allow_device_flow: bool,
101    /// Held across a cache refresh to prevent concurrent updates.
102    refresh_lock: Arc<Mutex<()>>,
103}
104
105pub use crate::providers::internal::auth::AuthError;
106
107#[derive(Debug, Clone)]
108pub struct AuthContext {
109    /// Resolved credential. Use [`Secret::expose`] only when raw bytes are required.
110    pub api_key: Secret,
111    pub api_base: Option<String>,
112}
113
114impl Authenticator {
115    pub fn new(
116        source: AuthSource,
117        access_token_file: Option<PathBuf>,
118        api_key_file: Option<PathBuf>,
119        device_code_handler: DeviceCodeHandler,
120        allow_device_flow: bool,
121    ) -> Self {
122        Self {
123            source,
124            access_token_file,
125            api_key_file,
126            device_code_handler,
127            allow_device_flow,
128            refresh_lock: Arc::default(),
129        }
130    }
131
132    /// Resolve the API key and optional API base, refreshing through `http` as needed.
133    /// Return cache, transport, or authorization errors. OAuth device login is
134    /// unsupported on WASM; access-token exchange remains available.
135    pub async fn auth_context<H>(&self, http: &H) -> Result<AuthContext, AuthError>
136    where
137        H: HttpClientExt,
138    {
139        match &self.source {
140            AuthSource::ApiKey(api_key) => Ok(AuthContext {
141                api_key: api_key.clone().into(),
142                api_base: None,
143            }),
144            #[cfg(not(target_family = "wasm"))]
145            AuthSource::GitHubAccessToken(access_token) => {
146                self.auth_context_with_github_access_token(http, access_token)
147                    .await
148            }
149            #[cfg(target_family = "wasm")]
150            AuthSource::GitHubAccessToken(access_token) => {
151                Ok(refresh_api_key(http, access_token).await?.into_context())
152            }
153            #[cfg(not(target_family = "wasm"))]
154            AuthSource::OAuth => self.auth_context_oauth(http).await,
155            #[cfg(target_family = "wasm")]
156            AuthSource::OAuth => Err(AuthError::Message(
157                "GitHub Copilot OAuth is not supported on wasm targets".into(),
158            )),
159        }
160    }
161}
162
163/// Copilot API key exchanged for a GitHub access token, and its native cache record.
164#[derive(Debug, Clone, Deserialize, Serialize, Default)]
165struct ApiKeyRecord {
166    token: Option<String>,
167    expires_at: Option<i64>,
168    endpoints: Option<ApiKeyEndpoints>,
169    bootstrap_token_fingerprint: Option<String>,
170}
171
172#[derive(Debug, Clone, Deserialize, Serialize, Default)]
173struct ApiKeyEndpoints {
174    api: Option<String>,
175}
176
177impl ApiKeyRecord {
178    fn api_base(&self) -> Option<String> {
179        self.endpoints
180            .as_ref()
181            .and_then(|endpoints| endpoints.api.as_ref())
182            .cloned()
183    }
184
185    fn into_context(self) -> AuthContext {
186        AuthContext {
187            api_base: self.api_base(),
188            api_key: self.token.unwrap_or_default().into(),
189        }
190    }
191}
192
193/// Exchange a GitHub access token for a Copilot API key.
194/// Return transport errors, or an error when the response has no non-blank token.
195async fn refresh_api_key<H>(http: &H, access_token: &str) -> Result<ApiKeyRecord, AuthError>
196where
197    H: HttpClientExt,
198{
199    let response: ApiKeyRecord = send_json(
200        http,
201        request(Method::GET, GITHUB_API_KEY_URL)
202            .header(http::header::ACCEPT, "application/json")
203            .header("editor-version", super::EDITOR_VERSION)
204            .header("editor-plugin-version", super::EDITOR_PLUGIN_VERSION)
205            .header("user-agent", super::USER_AGENT)
206            .header(http::header::AUTHORIZATION, format!("token {access_token}"))
207            .body(bytes::Bytes::new()),
208    )
209    .await?;
210
211    if response
212        .token
213        .as_ref()
214        .is_none_or(|token| token.trim().is_empty())
215    {
216        return Err(AuthError::Message(
217            "GitHub Copilot API key response did not include a token".into(),
218        ));
219    }
220
221    Ok(response)
222}