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#[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
49struct 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
72pub 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 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 #[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(¶ms)
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}"); 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 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()), "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()), "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()), "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()), "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()), "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()), "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()), "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}