Skip to main content

camel_auth/
oauth2.rs

1use async_trait::async_trait;
2use serde::Deserialize;
3use serde::de::Deserializer;
4use std::fmt;
5use std::time::{Duration, Instant};
6use tokio::sync::{Mutex, RwLock};
7use zeroize::Zeroizing;
8
9use camel_api::SsrfPolicy;
10
11use crate::http_client::{SsrfClientOptions, build_ssrf_pinned_client};
12use crate::types::AuthError;
13
14fn deserialize_zeroizing_string<'de, D>(deserializer: D) -> Result<Zeroizing<String>, D::Error>
15where
16    D: Deserializer<'de>,
17{
18    let s = String::deserialize(deserializer)?;
19    Ok(Zeroizing::new(s))
20}
21
22const DEFAULT_SKEW: Duration = Duration::from_secs(30);
23
24#[async_trait]
25pub trait TokenProvider: Send + Sync + std::fmt::Debug {
26    async fn get_token(&self) -> Result<String, AuthError>;
27}
28
29/// ADR-0051 credential boundary: manual-redaction
30#[derive(Deserialize)]
31struct TokenResponse {
32    #[serde(deserialize_with = "deserialize_zeroizing_string")]
33    access_token: Zeroizing<String>,
34    #[allow(dead_code)]
35    token_type: String,
36    expires_in: u64,
37}
38
39impl fmt::Debug for TokenResponse {
40    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
41        f.debug_struct("TokenResponse")
42            .field("access_token", &"[REDACTED]")
43            .field("token_type", &self.token_type)
44            .field("expires_in", &self.expires_in)
45            .finish()
46    }
47}
48
49/// ADR-0051 credential boundary: manual-redaction
50struct CachedToken {
51    access_token: Zeroizing<String>,
52    #[allow(dead_code)]
53    expires_at: Instant,
54    refresh_at: Instant,
55}
56
57impl CachedToken {
58    fn new(access_token: Zeroizing<String>, expires_in: Duration, skew: Duration) -> Self {
59        let expires_at = Instant::now() + expires_in;
60        Self {
61            access_token,
62            refresh_at: expires_at.checked_sub(skew).unwrap_or(expires_at),
63            expires_at,
64        }
65    }
66
67    fn is_usable(&self) -> bool {
68        Instant::now() < self.refresh_at
69    }
70}
71
72/// ADR-0051 credential boundary: manual-redaction
73pub struct ClientCredentialsProvider {
74    token_endpoint: String,
75    client_id: String,
76    client_secret: Zeroizing<String>,
77    scope: Option<String>,
78    audience: Option<Vec<String>>,
79    cache: RwLock<Option<CachedToken>>,
80    refresh_lock: Mutex<()>,
81    http: reqwest::Client,
82}
83
84impl std::fmt::Debug for ClientCredentialsProvider {
85    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
86        f.debug_struct("ClientCredentialsProvider")
87            .field("token_endpoint", &self.token_endpoint)
88            .field("client_id", &self.client_id)
89            .field("scope", &self.scope)
90            .field("audience", &self.audience)
91            .finish_non_exhaustive()
92    }
93}
94
95impl ClientCredentialsProvider {
96    /// Creates a production OAuth2 provider with HTTPS enforcement,
97    /// SSRF guard, DNS-rebinding protection, and hardened timeouts.
98    ///
99    /// Validates the token endpoint is a public HTTPS URI, resolves DNS
100    /// and pins validated IPs on the HTTP client — closing the TOCTOU
101    /// window between validation and the first outbound token request.
102    pub async fn new(
103        token_endpoint: String,
104        client_id: String,
105        client_secret: String,
106        scope: Option<String>,
107        audience: Option<Vec<String>>,
108        policy: SsrfPolicy,
109    ) -> Result<Self, AuthError> {
110        let http = build_ssrf_pinned_client(
111            &token_endpoint,
112            "OAuth2 token endpoint",
113            &SsrfClientOptions::new(policy)
114                .with_connect_timeout(Duration::from_secs(10))
115                .with_request_timeout(Duration::from_secs(30)),
116        )
117        .await?;
118        Ok(Self {
119            token_endpoint,
120            client_id,
121            client_secret: Zeroizing::new(client_secret),
122            scope,
123            audience,
124            cache: RwLock::new(None),
125            refresh_lock: Mutex::new(()),
126            http,
127        })
128    }
129
130    /// Test-only constructor that accepts a pre-built HTTP client.
131    ///
132    /// Skips SSRF validation and allows injecting a mock-capable client
133    /// (e.g. `wiremock`-configured). Do NOT use in production code.
134    #[doc(hidden)]
135    pub fn new_unchecked_for_test(
136        token_endpoint: String,
137        client_id: String,
138        client_secret: String,
139        scope: Option<String>,
140        audience: Option<Vec<String>>,
141        http: reqwest::Client,
142    ) -> Self {
143        Self {
144            token_endpoint,
145            client_id,
146            client_secret: Zeroizing::new(client_secret),
147            scope,
148            audience,
149            cache: RwLock::new(None),
150            refresh_lock: Mutex::new(()),
151            http,
152        }
153    }
154
155    async fn fetch_token(&self) -> Result<CachedToken, AuthError> {
156        let secret = self.client_secret.as_str();
157        let mut params: Vec<(&str, &str)> = vec![
158            ("grant_type", "client_credentials"),
159            ("client_id", &self.client_id),
160            ("client_secret", secret),
161        ];
162        if let Some(ref scope) = self.scope {
163            params.push(("scope", scope));
164        }
165        if let Some(ref audience) = self.audience {
166            for aud in audience {
167                params.push(("resource", aud));
168            }
169        }
170
171        let resp = self
172            .http
173            .post(&self.token_endpoint)
174            .form(&params)
175            .send()
176            .await
177            .map_err(|e| AuthError::ProviderUnavailable(format!("OAuth2 request failed: {e}")))?;
178
179        if !resp.status().is_success() {
180            let status = resp.status();
181            let body = resp.text().await.unwrap_or_default();
182            let sanitized = if body.len() > 128 {
183                format!("{}...(truncated)", &body[..128])
184            } else {
185                body
186            };
187            let message = format!("token endpoint returned {status}: {sanitized}"); // allow-secret
188            return Err(AuthError::ProviderUnavailable(message));
189        }
190
191        let token_resp: TokenResponse = resp
192            .json()
193            .await
194            .map_err(|e| AuthError::ProviderUnavailable(format!("invalid OAuth2 response: {e}")))?;
195
196        Ok(CachedToken::new(
197            token_resp.access_token,
198            Duration::from_secs(token_resp.expires_in),
199            DEFAULT_SKEW,
200        ))
201    }
202}
203
204#[async_trait]
205impl TokenProvider for ClientCredentialsProvider {
206    async fn get_token(&self) -> Result<String, AuthError> {
207        {
208            let cache = self.cache.read().await;
209            if let Some(ref cached) = *cache
210                && cached.is_usable()
211            {
212                return Ok(cached.access_token.as_str().to_owned());
213            }
214        }
215
216        let _guard = self.refresh_lock.lock().await;
217
218        {
219            let cache = self.cache.read().await;
220            if let Some(ref cached) = *cache
221                && cached.is_usable()
222            {
223                return Ok(cached.access_token.as_str().to_owned());
224            }
225        }
226
227        let cached = self.fetch_token().await?;
228        let token = cached.access_token.as_str().to_owned();
229        {
230            let mut cache = self.cache.write().await;
231            *cache = Some(cached);
232        }
233        Ok(token)
234    }
235}
236
237#[cfg(test)]
238mod tests {
239    use std::sync::Arc;
240
241    use super::*;
242    use wiremock::matchers::{body_string_contains, method, path};
243    use wiremock::{Mock, MockServer, ResponseTemplate};
244
245    #[tokio::test]
246    async fn oauth2_rejects_private_ip_token_endpoint() {
247        // validate_uri catches IP literals like
248        // 169.254.169.254 before DNS resolution
249        let result = ClientCredentialsProvider::new(
250            "https://169.254.169.254/token".into(),
251            "client".into(),
252            "secret".into(),
253            None,
254            None,
255            SsrfPolicy::PublicHttpsOnly,
256        )
257        .await;
258        assert!(
259            result.is_err(),
260            "private IP token endpoint should be rejected"
261        );
262    }
263
264    fn token_response(access_token: &str, expires_in: u64) -> serde_json::Value {
265        serde_json::json!({
266            "access_token": access_token,
267            "token_type": "Bearer",
268            "expires_in": expires_in,
269        })
270    }
271
272    #[tokio::test]
273    async fn test_get_token_fresh() {
274        let server = MockServer::start().await;
275        Mock::given(method("POST"))
276            .and(path("/protocol/openid-connect/token"))
277            .respond_with(ResponseTemplate::new(200).set_body_json(token_response("abc123", 300)))
278            .mount(&server)
279            .await;
280
281        let provider = ClientCredentialsProvider::new_unchecked_for_test(
282            format!("{}/protocol/openid-connect/token", server.uri()), // allow-secret
283            "test-client".into(),
284            "test-secret".into(),
285            None,
286            None,
287            reqwest::Client::new(),
288        );
289        let token = provider.get_token().await.unwrap();
290        assert_eq!(token, "abc123");
291    }
292
293    #[tokio::test]
294    async fn test_get_token_uses_cache() {
295        let server = MockServer::start().await;
296        Mock::given(method("POST"))
297            .respond_with(ResponseTemplate::new(200).set_body_json(token_response("cached", 300)))
298            .expect(1)
299            .mount(&server)
300            .await;
301
302        let provider = ClientCredentialsProvider::new_unchecked_for_test(
303            format!("{}/protocol/openid-connect/token", server.uri()), // allow-secret
304            "c".into(),
305            "s".into(),
306            None,
307            None,
308            reqwest::Client::new(),
309        );
310        let t1 = provider.get_token().await.unwrap();
311        let t2 = provider.get_token().await.unwrap();
312        assert_eq!(t1, "cached");
313        assert_eq!(t2, "cached");
314    }
315
316    #[tokio::test]
317    async fn test_get_token_refreshes_when_stale() {
318        let server = MockServer::start().await;
319        Mock::given(method("POST"))
320            .respond_with(ResponseTemplate::new(200).set_body_json(token_response("first", 1)))
321            .up_to_n_times(1)
322            .mount(&server)
323            .await;
324        Mock::given(method("POST"))
325            .respond_with(ResponseTemplate::new(200).set_body_json(token_response("second", 300)))
326            .mount(&server)
327            .await;
328
329        let provider = ClientCredentialsProvider::new_unchecked_for_test(
330            format!("{}/protocol/openid-connect/token", server.uri()), // allow-secret
331            "c".into(),
332            "s".into(),
333            None,
334            None,
335            reqwest::Client::new(),
336        );
337        let t1 = provider.get_token().await.unwrap();
338        assert_eq!(t1, "first");
339        tokio::time::sleep(Duration::from_millis(1100)).await;
340        let t2 = provider.get_token().await.unwrap();
341        assert_eq!(t2, "second");
342    }
343
344    #[tokio::test]
345    async fn test_get_token_server_error() {
346        let server = MockServer::start().await;
347        Mock::given(method("POST"))
348            .respond_with(ResponseTemplate::new(500))
349            .mount(&server)
350            .await;
351
352        let provider = ClientCredentialsProvider::new_unchecked_for_test(
353            format!("{}/protocol/openid-connect/token", server.uri()), // allow-secret
354            "c".into(),
355            "s".into(),
356            None,
357            None,
358            reqwest::Client::new(),
359        );
360        let err = provider.get_token().await.unwrap_err();
361        assert!(matches!(err, AuthError::ProviderUnavailable(_)));
362    }
363
364    #[tokio::test]
365    async fn test_get_token_invalid_response() {
366        let server = MockServer::start().await;
367        Mock::given(method("POST"))
368            .respond_with(
369                ResponseTemplate::new(200)
370                    .set_body_json(serde_json::json!({"error": "invalid_grant"})),
371            )
372            .mount(&server)
373            .await;
374
375        let provider = ClientCredentialsProvider::new_unchecked_for_test(
376            format!("{}/protocol/openid-connect/token", server.uri()), // allow-secret
377            "c".into(),
378            "s".into(),
379            None,
380            None,
381            reqwest::Client::new(),
382        );
383        let err = provider.get_token().await.unwrap_err();
384        assert!(matches!(err, AuthError::ProviderUnavailable(_)));
385    }
386
387    #[tokio::test]
388    async fn test_get_token_sends_audience_as_resource() {
389        let server = MockServer::start().await;
390        Mock::given(method("POST"))
391            .and(body_string_contains(
392                "resource=https%3A%2F%2Fapi.example.com",
393            ))
394            .respond_with(
395                ResponseTemplate::new(200).set_body_json(token_response("aud-token", 300)),
396            )
397            .mount(&server)
398            .await;
399
400        let provider = ClientCredentialsProvider::new_unchecked_for_test(
401            format!("{}/protocol/openid-connect/token", server.uri()), // allow-secret
402            "c".into(),
403            "s".into(),
404            None,
405            Some(vec!["https://api.example.com".into()]),
406            reqwest::Client::new(),
407        );
408        let token = provider.get_token().await.unwrap();
409        assert_eq!(token, "aud-token");
410    }
411
412    #[tokio::test]
413    async fn test_single_flight_concurrent_callers() {
414        let server = MockServer::start().await;
415        Mock::given(method("POST"))
416            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
417                "access_token": "single-flight-token",
418                "token_type": "Bearer",
419                "expires_in": 300,
420            })))
421            .expect(1)
422            .mount(&server)
423            .await;
424
425        let provider = Arc::new(ClientCredentialsProvider::new_unchecked_for_test(
426            format!("{}/protocol/openid-connect/token", server.uri()), // allow-secret
427            "c".into(),
428            "s".into(),
429            None,
430            None,
431            reqwest::Client::new(),
432        ));
433
434        let mut handles = vec![];
435        for _ in 0..5 {
436            let p = Arc::clone(&provider);
437            handles.push(tokio::spawn(async move { p.get_token().await }));
438        }
439        for h in handles {
440            let token = h.await.unwrap().unwrap();
441            assert_eq!(token, "single-flight-token");
442        }
443    }
444
445    #[test]
446    fn debug_redacts_access_token() {
447        let resp = TokenResponse {
448            access_token: Zeroizing::new("SENTINEL-OAUTH-TOKEN".to_string()),
449            token_type: "Bearer".to_string(),
450            expires_in: 300,
451        };
452        let debug = format!("{:?}", resp);
453        assert!(
454            !debug.contains("SENTINEL-OAUTH-TOKEN"),
455            "Debug output must not contain access_token: {debug}"
456        );
457        assert!(debug.contains("[REDACTED]"));
458    }
459}