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