Skip to main content

rig_core/providers/internal/
auth.rs

1//! Shared OAuth errors, device-code callbacks, and authentication transport helpers.
2//!
3//! ```
4//! use rig_core::providers::chatgpt::auth::DeviceCodeHandler;
5//! let handler = DeviceCodeHandler::new(|prompt| {
6//!     println!("Visit {} and enter {}", prompt.verification_uri, prompt.user_code);
7//! });
8//! ```
9
10use crate::http_client::{self, HttpClientExt};
11use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
12use std::sync::Arc;
13
14/// Device authorization details surfaced to a provider callback.
15#[derive(Debug, Clone)]
16pub struct DeviceCodePrompt {
17    /// URL where the user authorizes the device.
18    pub verification_uri: String,
19    /// Short code the user enters at the verification URL.
20    pub user_code: String,
21}
22
23/// Device-code callback with thread-safety bounds outside browser WASM.
24/// Browser WASM authenticators do not invoke it.
25#[cfg(not(all(target_arch = "wasm32", target_os = "unknown")))]
26pub(crate) type DeviceCodeCallback = dyn Fn(DeviceCodePrompt) + Send + Sync;
27#[cfg(all(target_arch = "wasm32", target_os = "unknown"))]
28pub(crate) type DeviceCodeCallback = dyn Fn(DeviceCodePrompt);
29
30/// Optional callback invoked when an OAuth device flow needs user action.
31#[derive(Clone, Default)]
32pub struct DeviceCodeHandler(pub(crate) Option<Arc<DeviceCodeCallback>>);
33
34impl DeviceCodeHandler {
35    /// Wraps a device-code callback.
36    pub fn new<F>(handler: F) -> Self
37    where
38        F: Fn(DeviceCodePrompt) + WasmCompatSend + WasmCompatSync + 'static,
39    {
40        Self(Some(Arc::new(handler)))
41    }
42}
43
44impl std::fmt::Debug for DeviceCodeHandler {
45    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
46        if self.0.is_some() {
47            f.write_str("DeviceCodeHandler(<callback>)")
48        } else {
49            f.write_str("DeviceCodeHandler(None)")
50        }
51    }
52}
53
54#[derive(Debug, thiserror::Error)]
55pub enum AuthError {
56    #[error("{0}")]
57    Message(String),
58    #[error(transparent)]
59    Io(#[from] std::io::Error),
60    #[error(transparent)]
61    Json(#[from] serde_json::Error),
62    /// The HTTP transport failed. Non-success responses arrive as the
63    /// status-bearing [`http_client::Error`] variants (so the status is still
64    /// inspectable); response-less failures as [`http_client::Error::Instance`].
65    #[error(transparent)]
66    Http(#[from] http_client::Error),
67}
68
69/// Build a request to an absolute authentication URL, independent of the API base.
70pub(crate) fn request(method: http::Method, url: &str) -> http::request::Builder {
71    http::Request::builder().method(method).uri(url)
72}
73
74/// Send `req` and decode its JSON response.
75/// Request and transport failures return [`AuthError::Http`], preserving HTTP
76/// status when available. Invalid JSON returns [`AuthError::Json`].
77pub(crate) async fn send_json<H, T>(
78    http: &H,
79    req: http::Result<http::Request<bytes::Bytes>>,
80) -> Result<T, AuthError>
81where
82    H: HttpClientExt,
83    T: serde::de::DeserializeOwned,
84{
85    let bytes = send_bytes(http, req).await?;
86    Ok(serde_json::from_slice(&bytes)?)
87}
88
89/// Send `req` through the transport and return the raw success body.
90pub(crate) async fn send_bytes<H>(
91    http: &H,
92    req: http::Result<http::Request<bytes::Bytes>>,
93) -> Result<bytes::Bytes, AuthError>
94where
95    H: HttpClientExt,
96{
97    let req = req.map_err(http_client::Error::Protocol)?;
98    let response = http.send::<_, bytes::Bytes>(req).await?;
99    Ok(response.into_body().await?)
100}
101
102/// Platform config directory used for on-disk OAuth/token caches
103/// (`APPDATA` on Windows; `XDG_CONFIG_HOME` falling back to `~/.config`
104/// elsewhere).
105pub(crate) fn config_dir() -> Option<std::path::PathBuf> {
106    use std::path::PathBuf;
107
108    #[cfg(target_os = "windows")]
109    {
110        std::env::var_os("APPDATA").map(PathBuf::from)
111    }
112
113    #[cfg(not(target_os = "windows"))]
114    {
115        std::env::var_os("XDG_CONFIG_HOME")
116            .map(PathBuf::from)
117            .or_else(|| std::env::var_os("HOME").map(|home| PathBuf::from(home).join(".config")))
118    }
119}
120
121/// Device-flow helpers for the native ChatGPT and Copilot authenticators.
122#[cfg(not(target_family = "wasm"))]
123pub(crate) mod device;