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