1use crate::error::AuthError;
16use serde::{Deserialize, Serialize};
17use std::sync::{Arc, LazyLock};
18use tokio::sync::{Mutex, RwLock};
19
20static HTTP_CLIENT: LazyLock<reqwest::Client> = LazyLock::new(|| {
22 reqwest::Client::builder()
23 .timeout(std::time::Duration::from_secs(5)) .pool_max_idle_per_host(10) .pool_idle_timeout(std::time::Duration::from_secs(30))
26 .build()
27 .expect("Failed to create HTTP client")
28});
29
30#[derive(Debug, Serialize, Deserialize, Clone)]
32pub struct Jwk {
33 pub kid: String,
35 pub kty: String,
37 pub alg: Option<String>,
39 #[serde(rename = "use")]
41 pub key_use: Option<String>,
42 pub key_ops: Option<Vec<String>>,
44 pub crv: Option<String>,
46 pub x: Option<String>,
48 pub y: Option<String>,
50 pub n: Option<String>,
52 pub e: Option<String>,
54 pub ext: Option<bool>,
56}
57
58#[derive(Debug, Serialize, Deserialize, Clone)]
60pub struct JwksResponse {
61 pub keys: Vec<Jwk>,
63}
64
65const JWKS_CACHE_DURATION: u64 = 24 * 3600; const JWKS_CACHE_MAX_AGE: u64 = 7 * 24 * 3600; #[derive(Debug, Clone)]
70pub struct JwksCache {
71 cache: Arc<RwLock<Option<JwksResponse>>>,
73 expires_at: Arc<RwLock<Option<u64>>>,
75 cached_at: Arc<RwLock<Option<u64>>>,
77 jwks_url: String,
79 fetch_mutex: Arc<Mutex<()>>,
81}
82
83impl JwksCache {
84 pub fn new(jwks_url: &str) -> Self {
90 if !jwks_url.starts_with("https://") {
92 tracing::warn!("JWKS URL should use HTTPS: {}", jwks_url);
93 }
94
95 Self {
96 cache: Arc::new(RwLock::new(None)),
97 expires_at: Arc::new(RwLock::new(None)),
98 cached_at: Arc::new(RwLock::new(None)),
99 jwks_url: jwks_url.to_string(),
100 fetch_mutex: Arc::new(Mutex::new(())),
101 }
102 }
103
104 pub async fn get_jwks(&self) -> Result<JwksResponse, AuthError> {
108 self.get_jwks_with_fallback().await
109 }
110
111 async fn get_jwks_with_fallback(&self) -> Result<JwksResponse, AuthError> {
113 if let Some(cached) = self.get_cached_jwks().await {
115 tracing::debug!("Using valid cached JWKS data");
116 return Ok(cached);
117 }
118
119 match self.fetch_fresh_jwks().await {
121 Ok(jwks) => {
122 tracing::info!("Successfully refreshed JWKS cache");
123 Ok(jwks)
124 }
125 Err(e) => {
126 tracing::warn!("Failed to refresh JWKS, attempting fallback: {:?}", e);
127 self.get_stale_cache().await
129 }
130 }
131 }
132
133 async fn get_cached_jwks(&self) -> Option<JwksResponse> {
135 let now = std::time::SystemTime::now()
136 .duration_since(std::time::UNIX_EPOCH)
137 .unwrap()
138 .as_secs();
139
140 let expires_at = *self.expires_at.read().await;
141 if let Some(expires) = expires_at {
142 if now < expires {
143 return self.cache.read().await.clone();
144 }
145 }
146 None
147 }
148
149 async fn get_stale_cache(&self) -> Result<JwksResponse, AuthError> {
151 let now = std::time::SystemTime::now()
152 .duration_since(std::time::UNIX_EPOCH)
153 .unwrap()
154 .as_secs();
155
156 let cached_at = *self.cached_at.read().await;
157 if let Some(cache_time) = cached_at {
158 if now - cache_time <= JWKS_CACHE_MAX_AGE {
159 if let Some(cached) = self.cache.read().await.clone() {
160 tracing::warn!(
161 "Using stale JWKS cache as fallback (age: {} hours)",
162 (now - cache_time) / 3600
163 );
164 return Ok(cached);
165 }
166 }
167 }
168
169 let error_msg = "No valid JWKS cache available and network fetch failed";
170 tracing::error!("{}", error_msg);
171 Err(AuthError::JwksError(error_msg.to_string()))
172 }
173
174 async fn fetch_fresh_jwks(&self) -> Result<JwksResponse, AuthError> {
176 let _fetch_guard = self.fetch_mutex.lock().await;
178
179 if let Some(cached) = self.get_cached_jwks().await {
181 tracing::debug!("JWKS cache was updated while waiting for lock");
182 return Ok(cached);
183 }
184
185 let now = std::time::SystemTime::now()
186 .duration_since(std::time::UNIX_EPOCH)
187 .unwrap()
188 .as_secs();
189
190 tracing::info!("Fetching fresh JWKS from: {}", self.jwks_url);
192
193 let response = HTTP_CLIENT.get(&self.jwks_url).send().await.map_err(|e| {
194 let error_msg = format!("Failed to fetch JWKS: {e:?}");
195 tracing::error!("{}", error_msg);
196 AuthError::JwksError(error_msg)
197 })?;
198
199 if !response.status().is_success() {
200 let error_msg = format!("JWKS endpoint returned status: {}", response.status());
201 tracing::error!("{}", error_msg);
202 return Err(AuthError::JwksError(error_msg));
203 }
204
205 let jwks: JwksResponse = response.json().await.map_err(|e| {
206 let error_msg = format!("Failed to parse JWKS response: {e:?}");
207 tracing::error!("{}", error_msg);
208 AuthError::JwksError(error_msg)
209 })?;
210
211 if jwks.keys.is_empty() {
213 let error_msg = "JWKS response contains no keys";
214 tracing::error!("{}", error_msg);
215 return Err(AuthError::JwksError(error_msg.to_string()));
216 }
217
218 *self.cache.write().await = Some(jwks.clone());
220 *self.expires_at.write().await = Some(now + JWKS_CACHE_DURATION);
221 *self.cached_at.write().await = Some(now);
222
223 tracing::info!(
224 "JWKS cache updated, expires at: {} (cached at: {})",
225 now + JWKS_CACHE_DURATION,
226 now
227 );
228 Ok(jwks)
229 }
230
231 pub async fn find_key(&self, kid: &str) -> Result<Jwk, AuthError> {
237 let jwks = self.get_jwks().await?;
238
239 jwks.keys
240 .iter() .find(|key| key.kid == kid)
242 .cloned() .ok_or_else(|| {
244 tracing::warn!("Key with kid '{}' not found in JWKS", kid);
245 AuthError::NoMatchingKey
246 })
247 }
248}
249
250