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