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
90pub 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}