1use crate::claims::AuthContext;
2use crate::error::VerifyError;
3use crate::keys::VerifyingKey;
4use crate::token::{JwksVerifier, TokenVerifier};
5use std::sync::Arc;
6use std::time::{Duration, Instant};
7use tokio::sync::RwLock;
8
9#[derive(Clone)]
11struct CachedJwks {
12 verifier: JwksVerifier,
13 fetched_at: Instant,
14}
15
16pub struct AsyncVerifier {
18 inner: VerifierInner,
19 jwks_url: Option<String>,
20 cache_duration: Duration,
21 cached_jwks: Arc<RwLock<Option<CachedJwks>>>,
22 issuer: String,
24 audiences: crate::AudienceSet,
26 require_origin: bool,
27}
28
29enum VerifierInner {
30 Static(TokenVerifier),
31 Jwks(JwksVerifier),
32}
33
34impl AsyncVerifier {
35 pub fn with_static_key(
37 key: VerifyingKey,
38 issuer: impl Into<String>,
39 audience: impl Into<String>,
40 ) -> Self {
41 let issuer_str = issuer.into();
42 let audience_str = audience.into();
43 Self {
44 inner: VerifierInner::Static(TokenVerifier::new(
45 key,
46 issuer_str.clone(),
47 audience_str.clone(),
48 )),
49 jwks_url: None,
50 cache_duration: Duration::from_secs(3600), cached_jwks: Arc::new(RwLock::new(None)),
52 issuer: issuer_str,
53 audiences: crate::AudienceSet::single(audience_str),
54 require_origin: false,
55 }
56 }
57
58 pub fn with_jwks(
60 jwks: crate::token::Jwks,
61 issuer: impl Into<String>,
62 audience: impl Into<String>,
63 ) -> Self {
64 let issuer_str = issuer.into();
65 let audience_str = audience.into();
66 Self {
67 inner: VerifierInner::Jwks(JwksVerifier::new(
68 jwks,
69 issuer_str.clone(),
70 audience_str.clone(),
71 )),
72 jwks_url: None,
73 cache_duration: Duration::from_secs(3600),
74 cached_jwks: Arc::new(RwLock::new(None)),
75 issuer: issuer_str,
76 audiences: crate::AudienceSet::single(audience_str),
77 require_origin: false,
78 }
79 }
80
81 #[cfg(feature = "jwks")]
83 pub fn with_jwks_url(
84 url: impl Into<String>,
85 issuer: impl Into<String>,
86 audience: impl Into<String>,
87 ) -> Self {
88 let issuer_str = issuer.into();
89 let audience_str = audience.into();
90 Self {
91 inner: VerifierInner::Static(TokenVerifier::new(
92 VerifyingKey::from_bytes(&[0u8; 32]).expect("zero key should be valid"),
93 issuer_str.clone(),
94 audience_str.clone(),
95 )),
96 jwks_url: Some(url.into()),
97 issuer: issuer_str,
98 audiences: crate::AudienceSet::single(audience_str),
99 cache_duration: Duration::from_secs(3600),
100 cached_jwks: Arc::new(RwLock::new(None)),
101 require_origin: false,
102 }
103 }
104
105 pub fn with_origin_validation(mut self) -> Self {
107 self.require_origin = true;
108 self.inner = match self.inner {
109 VerifierInner::Static(verifier) => {
110 VerifierInner::Static(verifier.with_origin_validation())
111 }
112 VerifierInner::Jwks(verifier) => VerifierInner::Jwks(verifier.with_origin_validation()),
113 };
114 self
115 }
116
117 pub fn with_cache_duration(mut self, duration: Duration) -> Self {
119 self.cache_duration = duration;
120 self
121 }
122
123 #[cfg(feature = "jwks")]
125 pub async fn verify(
126 &self,
127 token: &str,
128 expected_origin: Option<&str>,
129 expected_client_ip: Option<&str>,
130 ) -> Result<AuthContext, VerifyError> {
131 match &self.inner {
133 VerifierInner::Static(verifier) => {
134 verifier.verify(token, expected_origin, expected_client_ip)
135 }
136 VerifierInner::Jwks(verifier) => {
137 verifier.verify(token, expected_origin, expected_client_ip)
138 }
139 }
140 }
141
142 #[cfg(not(feature = "jwks"))]
144 pub fn verify(
145 &self,
146 token: &str,
147 expected_origin: Option<&str>,
148 expected_client_ip: Option<&str>,
149 ) -> Result<AuthContext, VerifyError> {
150 match &self.inner {
151 VerifierInner::Static(verifier) => {
152 verifier.verify(token, expected_origin, expected_client_ip)
153 }
154 VerifierInner::Jwks(verifier) => {
155 verifier.verify(token, expected_origin, expected_client_ip)
156 }
157 }
158 }
159
160 #[cfg(feature = "jwks")]
162 pub async fn refresh_cache(&self) -> Result<(), VerifyError> {
163 if let Some(ref jwks_url) = self.jwks_url {
164 let jwks = crate::token::JwksVerifier::fetch_jwks(jwks_url)
166 .await
167 .map_err(|e| VerifyError::InvalidFormat(format!("Failed to fetch JWKS: {}", e)))?;
168
169 let verifier =
171 JwksVerifier::with_audience_set(jwks, &self.issuer, self.audiences.clone());
172 let verifier = if self.require_origin {
173 verifier.with_origin_validation()
174 } else {
175 verifier
176 };
177
178 let mut cached = self.cached_jwks.write().await;
180 *cached = Some(CachedJwks {
181 verifier,
182 fetched_at: Instant::now(),
183 });
184 }
185 Ok(())
186 }
187
188 async fn get_cached_verifier(&self) -> Option<JwksVerifier> {
190 let cached = self.cached_jwks.read().await;
191 if let Some(ref cached_jwks) = *cached {
192 if cached_jwks.fetched_at.elapsed() < self.cache_duration {
193 return Some(cached_jwks.verifier.clone());
194 }
195 }
196 None
197 }
198
199 #[cfg(feature = "jwks")]
201 pub async fn verify_with_cache(
202 &self,
203 token: &str,
204 expected_origin: Option<&str>,
205 expected_client_ip: Option<&str>,
206 ) -> Result<AuthContext, VerifyError> {
207 if let Some(verifier) = self.get_cached_verifier().await {
209 match verifier.verify(token, expected_origin, expected_client_ip) {
210 Ok(ctx) => return Ok(ctx),
211 Err(VerifyError::KeyNotFound(_)) => {
212 }
214 Err(e) => return Err(e),
215 }
216 }
217
218 self.refresh_cache().await?;
220
221 if let Some(verifier) = self.get_cached_verifier().await {
222 verifier.verify(token, expected_origin, expected_client_ip)
223 } else if self.jwks_url.is_some() {
224 Err(VerifyError::InvalidFormat(
225 "JWKS cache unavailable after refresh".to_string(),
226 ))
227 } else {
228 match &self.inner {
230 VerifierInner::Static(verifier) => {
231 verifier.verify(token, expected_origin, expected_client_ip)
232 }
233 VerifierInner::Jwks(verifier) => {
234 verifier.verify(token, expected_origin, expected_client_ip)
235 }
236 }
237 }
238 }
239}
240
241pub struct SimpleVerifier {
243 inner: TokenVerifier,
244}
245
246impl SimpleVerifier {
247 pub fn new(key: VerifyingKey, issuer: impl Into<String>, audience: impl Into<String>) -> Self {
249 Self {
250 inner: TokenVerifier::new(key, issuer, audience),
251 }
252 }
253
254 pub fn verify(
256 &self,
257 token: &str,
258 expected_origin: Option<&str>,
259 expected_client_ip: Option<&str>,
260 ) -> Result<AuthContext, VerifyError> {
261 self.inner
262 .verify(token, expected_origin, expected_client_ip)
263 }
264}
265
266#[cfg(test)]
267mod tests {
268 use super::*;
269 use crate::claims::{KeyClass, SessionClaims};
270 use crate::keys::SigningKey;
271 use crate::token::TokenSigner;
272 use base64::Engine;
273
274 #[cfg(feature = "jwks")]
275 use tokio::io::{AsyncReadExt, AsyncWriteExt};
276
277 #[tokio::test]
278 async fn test_async_verifier_with_static_key() {
279 let signing_key = SigningKey::generate();
280 let verifying_key = signing_key.verifying_key();
281
282 let signer = TokenSigner::new(signing_key, "test-issuer");
283 let verifier =
284 AsyncVerifier::with_static_key(verifying_key, "test-issuer", "test-audience");
285
286 let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
287 .with_scope("read")
288 .with_metering_key("meter-123")
289 .with_key_class(KeyClass::Publishable)
290 .build();
291
292 let token = signer.sign(claims).unwrap();
293 let context = verifier.verify(&token, None, None).await.unwrap();
294
295 assert_eq!(context.subject, "test-subject");
296 }
297
298 #[test]
299 fn test_simple_verifier() {
300 let signing_key = SigningKey::generate();
301 let verifying_key = signing_key.verifying_key();
302
303 let signer = TokenSigner::new(signing_key, "test-issuer");
304 let verifier = SimpleVerifier::new(verifying_key, "test-issuer", "test-audience");
305
306 let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
307 .with_scope("read")
308 .with_metering_key("meter-123")
309 .with_key_class(KeyClass::Publishable)
310 .build();
311
312 let token = signer.sign(claims).unwrap();
313 let context = verifier.verify(&token, None, None).unwrap();
314
315 assert_eq!(context.subject, "test-subject");
316 assert_eq!(context.metering_key, "meter-123");
317 }
318
319 #[cfg(feature = "jwks")]
320 #[test]
321 fn test_verify_with_cache_returns_explicit_error_when_cache_stays_empty() {
322 tokio::runtime::Runtime::new().unwrap().block_on(async {
323 let signing_key = SigningKey::generate();
324 let verifying_key = signing_key.verifying_key();
325 let signer = TokenSigner::new(signing_key, "test-issuer");
326
327 let jwks = serde_json::json!({
328 "keys": [{
329 "kty": "OKP",
330 "use": "sig",
331 "kid": verifying_key.key_id(),
332 "x": base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(verifying_key.to_bytes()),
333 }]
334 })
335 .to_string();
336
337 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
338 let addr = listener.local_addr().unwrap();
339 let response_body = jwks.clone();
340 tokio::spawn(async move {
341 let (mut socket, _) = listener.accept().await.unwrap();
342 let mut buffer = [0u8; 1024];
343 let _ = socket.read(&mut buffer).await;
344
345 let response = format!(
346 "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
347 response_body.len(),
348 response_body
349 );
350 socket.write_all(response.as_bytes()).await.unwrap();
351 });
352
353 let verifier = AsyncVerifier::with_jwks_url(
354 format!("http://{addr}/jwks"),
355 "test-issuer",
356 "test-audience",
357 )
358 .with_cache_duration(Duration::ZERO);
359
360 let claims = SessionClaims::builder("test-issuer", "test-subject", "test-audience")
361 .with_scope("read")
362 .with_metering_key("meter-123")
363 .with_key_class(KeyClass::Publishable)
364 .build();
365 let token = signer.sign(claims).unwrap();
366
367 let result = verifier.verify_with_cache(&token, None, None).await;
368 assert!(matches!(
369 result,
370 Err(VerifyError::InvalidFormat(ref msg)) if msg == "JWKS cache unavailable after refresh"
371 ));
372 });
373 }
374}