rig_core/providers/copilot/auth/
mod.rs1use crate::http_client::HttpClientExt;
2use crate::wire::Secret;
3use futures::lock::Mutex;
4use std::fmt;
5use std::path::PathBuf;
6use std::sync::Arc;
7
8pub use crate::providers::internal::auth::{DeviceCodeHandler, DeviceCodePrompt};
9
10#[cfg(not(target_family = "wasm"))]
11mod native;
12#[cfg(target_family = "wasm")]
13mod wasm;
14
15#[cfg(not(target_family = "wasm"))]
16use native as platform;
17#[cfg(target_family = "wasm")]
18use wasm as platform;
19
20pub fn default_token_dir() -> Option<PathBuf> {
23 crate::providers::internal::auth::config_dir().map(|dir| dir.join("github_copilot"))
24}
25
26#[derive(Clone)]
27pub enum AuthSource {
28 ApiKey(String),
29 GitHubAccessToken(String),
30 OAuth,
31}
32
33impl fmt::Debug for AuthSource {
34 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
35 match self {
36 Self::ApiKey(_) => f.write_str("ApiKey(<redacted>)"),
37 Self::GitHubAccessToken(_) => f.write_str("GitHubAccessToken(<redacted>)"),
38 Self::OAuth => f.write_str("OAuth"),
39 }
40 }
41}
42
43#[derive(Clone)]
44pub struct Authenticator {
45 source: AuthSource,
46 platform: Arc<Mutex<platform::PlatformAuthenticator>>,
48}
49
50impl fmt::Debug for Authenticator {
51 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
52 f.debug_struct("Authenticator")
53 .field("source", &self.source)
54 .field("platform", &"<serialized>")
55 .finish()
56 }
57}
58
59pub use crate::providers::internal::auth::AuthError;
60
61#[derive(Debug, Clone)]
62pub struct AuthContext {
63 pub api_key: Secret,
65 pub api_base: Option<String>,
66}
67
68impl Authenticator {
69 pub fn new(
70 source: AuthSource,
71 access_token_file: Option<PathBuf>,
72 api_key_file: Option<PathBuf>,
73 device_code_handler: DeviceCodeHandler,
74 allow_device_flow: bool,
75 ) -> Self {
76 Self {
77 source,
78 platform: Arc::new(Mutex::new(platform::PlatformAuthenticator::new(
79 access_token_file,
80 api_key_file,
81 device_code_handler,
82 allow_device_flow,
83 ))),
84 }
85 }
86
87 pub async fn auth_context<H>(&self, http: &H) -> Result<AuthContext, AuthError>
91 where
92 H: HttpClientExt,
93 {
94 match &self.source {
95 AuthSource::ApiKey(api_key) => Ok(AuthContext {
96 api_key: api_key.clone().into(),
97 api_base: None,
98 }),
99 AuthSource::GitHubAccessToken(access_token) => {
100 self.platform
101 .lock()
102 .await
103 .auth_context_with_github_access_token(http, access_token)
104 .await
105 }
106 AuthSource::OAuth => self.platform.lock().await.auth_context_oauth(http).await,
107 }
108 }
109}