rig_core/providers/copilot/auth/
mod.rs1use 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
18pub fn default_token_dir() -> Option<PathBuf> {
21 crate::providers::internal::auth::config_dir().map(|dir| dir.join("github_copilot"))
22}
23
24pub(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 refresh_lock: Arc<Mutex<()>>,
103}
104
105pub use crate::providers::internal::auth::AuthError;
106
107#[derive(Debug, Clone)]
108pub struct AuthContext {
109 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 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#[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
193async 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}