Skip to main content

atman_runtime/
oauth.rs

1use std::sync::Arc;
2use std::time::{SystemTime, UNIX_EPOCH};
3
4use anyhow::{Context, Result};
5use base64::Engine;
6use base64::engine::general_purpose::URL_SAFE_NO_PAD;
7use rand::RngCore;
8use sha2::{Digest, Sha256};
9
10use crate::auth_store::StoredProvider;
11use crate::config_hub::AuthTokenUpdate;
12use crate::provider::{DiscoveredModel, Provider};
13
14pub struct Pkce {
15    pub verifier: String,
16    pub challenge: String,
17}
18
19impl Pkce {
20    pub fn generate() -> Self {
21        let mut bytes = [0u8; 32];
22        rand::thread_rng().fill_bytes(&mut bytes);
23        let verifier = URL_SAFE_NO_PAD.encode(bytes);
24
25        let mut hasher = Sha256::new();
26        hasher.update(verifier.as_bytes());
27        let digest = hasher.finalize();
28        let challenge = URL_SAFE_NO_PAD.encode(digest);
29
30        Pkce {
31            verifier,
32            challenge,
33        }
34    }
35}
36
37pub struct TokenResult {
38    pub access_token: String,
39    pub refresh_token: Option<String>,
40    pub expires_at: i64,
41    pub account: Option<String>,
42}
43
44pub trait OAuthProvider: Provider {
45    fn authorize_url() -> (String, Pkce, String);
46    fn exchange_code(
47        code: &str,
48        verifier: &str,
49    ) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<TokenResult>> + Send>>;
50    fn refresh_token(
51        token: &str,
52    ) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<TokenResult>> + Send>>;
53    fn from_stored(stored: &StoredProvider) -> Self;
54}
55
56pub fn generate_state() -> String {
57    let mut bytes = [0u8; 16];
58    rand::thread_rng().fill_bytes(&mut bytes);
59    bytes.iter().map(|b| format!("{b:02x}")).collect()
60}
61
62pub fn parse_jwt_exp(token: &str) -> Option<i64> {
63    let parts: Vec<&str> = token.split('.').collect();
64    if parts.len() != 3 {
65        return None;
66    }
67    let payload = URL_SAFE_NO_PAD.decode(parts[1].as_bytes()).ok()?;
68    let v: serde_json::Value = serde_json::from_slice(&payload).ok()?;
69    v.get("exp")?.as_i64()
70}
71
72pub fn extract_account_from_id_token(id_token: &str) -> Option<String> {
73    let parts: Vec<&str> = id_token.split('.').collect();
74    if parts.len() != 3 {
75        return None;
76    }
77    let payload = URL_SAFE_NO_PAD.decode(parts[1].as_bytes()).ok()?;
78    let v: serde_json::Value = serde_json::from_slice(&payload).ok()?;
79    v.get("email")
80        .and_then(|e| e.as_str())
81        .map(|s| s.to_string())
82}
83
84pub async fn create_oauth_provider<P: OAuthProvider>(
85    stored: &StoredProvider,
86) -> Result<(Arc<P>, Vec<DiscoveredModel>)> {
87    create_oauth_provider_impl(stored, true).await
88}
89
90/// Same as `create_oauth_provider` but skips `discover_models()`.
91/// Returns an empty model list. Use when discovery will happen in the background.
92pub async fn create_oauth_provider_no_discover<P: OAuthProvider>(
93    stored: &StoredProvider,
94) -> Result<Arc<P>> {
95    let (provider, _) = create_oauth_provider_impl(stored, false).await?;
96    Ok(provider)
97}
98
99async fn create_oauth_provider_impl<P: OAuthProvider>(
100    stored: &StoredProvider,
101    discover: bool,
102) -> Result<(Arc<P>, Vec<DiscoveredModel>)> {
103    let now = SystemTime::now()
104        .duration_since(UNIX_EPOCH)
105        .context("system clock")?
106        .as_secs() as i64;
107    let refresh_window = 300;
108
109    let mut updated = stored.clone();
110    if now + refresh_window >= stored.expires_at {
111        if let Some(ref rt) = stored.refresh_token {
112            let tokens = P::refresh_token(rt).await?;
113            let persisted_refresh_token = tokens.refresh_token.clone();
114            let persisted_account = tokens.account.clone();
115            updated.access_token = tokens.access_token;
116            updated.expires_at = tokens.expires_at;
117            if tokens.refresh_token.is_some() {
118                updated.refresh_token = tokens.refresh_token;
119            }
120            let _ = crate::config_hub::ConfigHub::global().and_then(|hub| {
121                hub.update_auth_tokens(
122                    &stored.id,
123                    AuthTokenUpdate {
124                        access_token: updated.access_token.clone(),
125                        refresh_token: persisted_refresh_token,
126                        expires_at: updated.expires_at,
127                        account: persisted_account,
128                    },
129                )
130                .map(|_| ())
131            });
132        }
133    }
134
135    let provider = P::from_stored(&updated);
136    let models = if discover {
137        provider.discover_models().await
138    } else {
139        vec![]
140    };
141    Ok((Arc::new(provider), models))
142}
143
144pub fn callback_page(ok: bool, title: &str, message: &str) -> String {
145    let icon = if ok { "✓" } else { "✗" };
146    let color = if ok { "#0078a0" } else { "#c0392b" };
147    let acetate = if ok {
148        "rgba(0,120,160,0.10)"
149    } else {
150        "rgba(192,57,43,0.10)"
151    };
152    format!(
153        r#"<!DOCTYPE html>
154<html lang="zh">
155<head>
156<meta charset="utf-8">
157<meta name="viewport" content="width=device-width,initial-scale=1">
158<title>{title} — atman</title>
159<style>
160  * {{ margin:0; padding:0; box-sizing:border-box; }}
161  body {{
162    font-family: "JetBrains Mono","Fira Code",Menlo,Consolas,monospace;
163    background: linear-gradient(180deg,#f0f0f0 0%,#e8e8ec 100%);
164    min-height: 100vh; display:flex; align-items:center; justify-content:center;
165  }}
166  .card {{
167    background: #fff; border-radius: 12px; padding: 36px 44px;
168    box-shadow: 0 2px 8px rgba(0,0,0,.06);
169    text-align: center; max-width: 520px;
170    border-top: 3px solid {color};
171  }}
172  .logo {{ margin-bottom: 20px; }}
173  .logo pre {{
174    font-size: 6.5px; line-height: 1.15; color: #0078a0;
175    font-family: "JetBrains Mono","Fira Code",Menlo,Consolas,monospace;
176  }}
177  .icon {{
178    font-size: 36px; color: {color}; margin-bottom: 16px;
179    display: inline-block; width: 56px; height: 56px; line-height: 56px;
180    border-radius: 50%; background: {acetate};
181  }}
182  h1 {{ font-size: 18px; font-weight: 600; color: #1e1e1e; margin-bottom: 8px; }}
183  p  {{ font-size: 13px; color: #606060; line-height: 1.6; }}
184</style>
185</head>
186<body>
187<div class="card">
188  <div class="logo"><pre>
189      ⢀⡤⣾⢿⡿⢿⡿⣷⢤⡀                                           
190     ⢠⢯⢎⠞⡵⠚⠓⢮⠳⡱⡽⡄                                          
191     ⡟⡏⡏⣀⣳⣀⣀⣞⣀⡰⢹⢻    ████████╗███╗   ███╗ █████╗ ███╗   ██╗
192  ⢀⣠⡄⣧⣇⡇⠻⠿⠿⠿⠿⠿⢿⡿⣷⣦⣄⡀ ╚══██╔══╝████╗ ████║██╔══██╗████╗  ██║
193⢀⡴⡫⡪⠕⠹⡼⡜⡄    ⢠⢢⢮⠍⠺⢗⢝⢦⡀  ██║   ██╔████╔██║███████║██╔██╗ ██║
194⡞⡞⡞   ⠙⣝⢞⢦⡀⢀⡴⡳⣫⠋   ⢳⢳⢳  ██║   ██║╚██╔╝██║██╔══██║██║╚██╗██║
195⢧⢧⡣⡀   ⠈⣓⡡⣔⣽⡪⢞⠁   ⢀⢜⡼⡼  ██║   ██║     ██║██║  ██║██║ ╚████║
196⠈⠓⠿⣾⣿⣿⣿⣿⡿⠿⠛⠙⠾⢷⣿⣿⣿⣿⣷⠿⠚⠁  ╚═╝   ╚═╝     ╚═╝╚═╝  ╚═╝╚═╝  ╚═══╝
197</pre></div>
198  <div class="icon">{icon}</div>
199  <h1>{title}</h1>
200  <p>{message}</p>
201</div>
202</body>
203</html>"#
204    )
205}
206
207#[cfg(test)]
208mod tests {
209    use super::*;
210    use base64::Engine;
211    use base64::engine::general_purpose::URL_SAFE_NO_PAD;
212
213    #[test]
214    fn pkce_generates_43char_verifier_and_valid_challenge() {
215        let pkce = Pkce::generate();
216        assert_eq!(pkce.verifier.len(), 43);
217        assert!(!pkce.challenge.is_empty());
218        let mut hasher = Sha256::new();
219        hasher.update(pkce.verifier.as_bytes());
220        let digest = hasher.finalize();
221        let expected = URL_SAFE_NO_PAD.encode(digest);
222        assert_eq!(pkce.challenge, expected);
223    }
224
225    #[test]
226    fn generate_state_is_32_hex_chars() {
227        let state = generate_state();
228        assert_eq!(state.len(), 32);
229        assert!(state.chars().all(|c| c.is_ascii_hexdigit()));
230    }
231
232    #[test]
233    fn parse_jwt_exp_works() {
234        let payload = URL_SAFE_NO_PAD.encode(r#"{"exp":123456789,"other":"data"}"#.as_bytes());
235        let token = format!("header.{payload}.sig");
236        assert_eq!(parse_jwt_exp(&token), Some(123456789));
237    }
238
239    #[test]
240    fn parse_jwt_exp_returns_none_when_missing() {
241        assert_eq!(parse_jwt_exp("not.a.jwt"), None);
242        assert_eq!(parse_jwt_exp(""), None);
243    }
244
245    #[test]
246    fn extract_account_prefers_email() {
247        let payload = URL_SAFE_NO_PAD.encode(r#"{"email":"a@b.com","sub":"123"}"#.as_bytes());
248        let token = format!("h.{payload}.sig");
249        assert_eq!(
250            extract_account_from_id_token(&token),
251            Some("a@b.com".to_string())
252        );
253    }
254
255    #[test]
256    fn extract_account_returns_none_when_missing() {
257        let payload = URL_SAFE_NO_PAD.encode(r#"{"name":"John"}"#.as_bytes());
258        let token = format!("h.{payload}.sig");
259        assert_eq!(extract_account_from_id_token(&token), None);
260    }
261
262    #[test]
263    fn callback_page_ok_has_right_icon_and_color() {
264        let html = callback_page(true, "OK", "done");
265        assert!(html.contains("✓"));
266        assert!(html.contains("#0078a0"));
267    }
268
269    #[test]
270    fn callback_page_err_has_right_icon_and_color() {
271        let html = callback_page(false, "Err", "fail");
272        assert!(html.contains("✗"));
273        assert!(html.contains("#c0392b"));
274    }
275}