Skip to main content

arete_auth/
verifier.rs

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/// Cached JWKS with expiration
10#[derive(Clone)]
11struct CachedJwks {
12    verifier: JwksVerifier,
13    fetched_at: Instant,
14}
15
16/// Async verifier with JWKS caching support
17pub struct AsyncVerifier {
18    inner: VerifierInner,
19    jwks_url: Option<String>,
20    cache_duration: Duration,
21    cached_jwks: Arc<RwLock<Option<CachedJwks>>>,
22    /// Issuer for JWKS-based verification
23    issuer: String,
24    /// Audience for JWKS-based verification
25    audiences: crate::AudienceSet,
26    require_origin: bool,
27}
28
29enum VerifierInner {
30    Static(TokenVerifier),
31    Jwks(JwksVerifier),
32}
33
34impl AsyncVerifier {
35    /// Create a verifier with a static key
36    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), // 1 hour default
51            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    /// Create a verifier with JWKS
59    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    /// Create a verifier that fetches JWKS from a URL
82    #[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    /// Require origin validation on verified tokens.
106    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    /// Set cache duration for JWKS
118    pub fn with_cache_duration(mut self, duration: Duration) -> Self {
119        self.cache_duration = duration;
120        self
121    }
122
123    /// Verify a token with automatic JWKS fetching and caching
124    #[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        // If using static JWKS or static key, use directly
132        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    /// Verify a token (non-JWKS version)
143    #[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    /// Refresh JWKS cache from the configured URL
161    #[cfg(feature = "jwks")]
162    pub async fn refresh_cache(&self) -> Result<(), VerifyError> {
163        if let Some(ref jwks_url) = self.jwks_url {
164            // Fetch JWKS from URL
165            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            // Create new verifier with fetched JWKS
170            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            // Update cache
179            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    /// Get cached verifier if available and not expired
189    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    /// Verify a token with automatic JWKS caching
200    #[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        // Try cached verifier first
208        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                    // Key not found in cache, refresh and retry
213                }
214                Err(e) => return Err(e),
215            }
216        }
217
218        // Refresh cache and try again
219        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            // Fallback to inner verifier if no cache available
229            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
241/// Simple synchronous verifier for use in non-async contexts
242pub struct SimpleVerifier {
243    inner: TokenVerifier,
244}
245
246impl SimpleVerifier {
247    /// Create a new simple verifier
248    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    /// Verify a token synchronously
255    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}