1use crate::clock::{Clock, SystemClock};
6use oauth2::basic::BasicClient;
7use oauth2::{AuthType, reqwest};
8use oauth2::{
9 AuthUrl, ClientId, ClientSecret, EndpointNotSet, EndpointSet, Scope, TokenResponse, TokenUrl,
10};
11use thiserror::Error;
12use url;
13
14const DEFAULT_REFRESH_SKEW_SECONDS: u64 = 30;
15
16#[derive(Debug, Error)]
17pub enum OAuth2TokenProviderError {
18 #[error("invalid config: {0}")]
19 InvalidConfig(&'static str),
20
21 #[error("invalid token_url: {0}")]
22 InvalidTokenUrl(#[from] url::ParseError),
23
24 #[error("failed to build http client: {0}")]
25 HttpClientBuild(String),
26
27 #[error("token request failed: {0}")]
28 TokenRequest(String),
29
30 #[error("token response missing expires_in")]
31 MissingExpiresIn,
32}
33
34#[derive(Debug, Clone)]
35pub struct ClientCredentialsConfig {
36 pub token_url: String,
37 pub client_id: String,
38 pub client_secret: String,
39 pub scopes: Vec<String>,
40 pub refresh_skew_seconds: Option<u64>,
41 pub accept_invalid_certs: Option<bool>,
42}
43
44#[derive(Debug, Clone)]
45struct CachedToken {
46 access_token: String,
47 expired_at: u64,
48}
49
50type ConfiguredBasicClient =
51 BasicClient<EndpointSet, EndpointNotSet, EndpointNotSet, EndpointNotSet, EndpointSet>;
52
53#[derive(Debug)]
54pub struct OAuth2TokenProvider<C: Clock = SystemClock> {
55 config: ClientCredentialsConfig,
56 client: ConfiguredBasicClient,
57 cache: tokio::sync::RwLock<Option<CachedToken>>,
58 clock: C,
59}
60
61pub type Result<T> = std::result::Result<T, OAuth2TokenProviderError>;
62
63impl OAuth2TokenProvider<SystemClock> {
64 pub fn new(config: ClientCredentialsConfig) -> Result<Self> {
65 Self::with_clock(config, SystemClock)
66 }
67}
68
69impl<C: Clock> OAuth2TokenProvider<C> {
70 pub fn with_clock(config: ClientCredentialsConfig, clock: C) -> Result<Self> {
71 Self::validate_config(&config)?;
72
73 let token_url = TokenUrl::new(config.token_url.clone())?;
74 let auth_url = AuthUrl::new("https://invalid.local/authorize".to_string())
75 .expect("hardcoded url must be valid");
76
77 let client = BasicClient::new(ClientId::new(config.client_id.clone()))
78 .set_client_secret(ClientSecret::new(config.client_secret.clone()))
79 .set_auth_uri(auth_url)
80 .set_token_uri(token_url)
81 .set_auth_type(AuthType::RequestBody);
82
83 Ok(Self {
84 config,
85 client,
86 cache: tokio::sync::RwLock::new(None),
87 clock,
88 })
89 }
90
91 fn validate_config(config: &ClientCredentialsConfig) -> Result<()> {
92 if config.token_url.trim().is_empty() {
93 return Err(OAuth2TokenProviderError::InvalidConfig(
94 "token_url must not be empty",
95 ));
96 }
97 if config.client_id.trim().is_empty() {
98 return Err(OAuth2TokenProviderError::InvalidConfig(
99 "client_id must not be empty",
100 ));
101 }
102 if config.client_secret.trim().is_empty() {
103 return Err(OAuth2TokenProviderError::InvalidConfig(
104 "client_secret must not be empty",
105 ));
106 }
107 Ok(())
108 }
109
110 pub fn refresh_skew(&self) -> u64 {
111 match self.config.refresh_skew_seconds {
112 Some(0) => 0,
113 Some(seconds) => seconds,
114 None => DEFAULT_REFRESH_SKEW_SECONDS,
115 }
116 }
117
118 pub fn get_token(&self) -> Result<String> {
119 let guard = self.cache.blocking_read();
120 if let Some(token) = guard.as_ref() {
121 let now = self.clock.now();
122 let skew = self.refresh_skew();
123 if skew == 0 || token.expired_at > now.saturating_add(skew) {
124 return Ok(token.access_token.clone());
125 }
126 }
127 drop(guard);
128
129 let http_client = reqwest::blocking::ClientBuilder::new()
130 .danger_accept_invalid_certs(self.config.accept_invalid_certs.unwrap_or(false)) .redirect(reqwest::redirect::Policy::none())
132 .build()
133 .map_err(|e| OAuth2TokenProviderError::HttpClientBuild(e.to_string()))?;
134
135 let mut req = self.client.exchange_client_credentials();
136
137 for s in &self.config.scopes {
138 let scope = s.trim();
139 if !scope.is_empty() {
140 req = req.add_scope(Scope::new(scope.to_string()));
141 }
142 }
143
144 let token = req
145 .request(&http_client)
146 .map_err(|e| OAuth2TokenProviderError::TokenRequest(e.to_string()))?;
147
148 let access_token = token.access_token().secret().to_string();
149 let expires_in = token
150 .expires_in()
151 .ok_or(OAuth2TokenProviderError::MissingExpiresIn)?
152 .as_secs();
153
154 let expired_at = self.clock.now().saturating_add(expires_in);
155
156 let mut guard = self.cache.blocking_write();
157 *guard = Some(CachedToken {
158 access_token: access_token.clone(),
159 expired_at,
160 });
161 drop(guard);
162
163 Ok(access_token)
164 }
165}
166
167#[cfg(test)]
168mod tests {
169 use super::*;
170 use crate::clock::ManualClock;
171 use httpmock::prelude::*;
172 use std::sync::Arc;
173
174 fn cfg(token_url: String) -> ClientCredentialsConfig {
175 ClientCredentialsConfig {
176 token_url,
177 client_id: "client-id".to_string(),
178 client_secret: "client-secret".to_string(),
179 scopes: vec!["scope-a".into(), "scope-b".into()],
180 refresh_skew_seconds: Some(30),
181 accept_invalid_certs: Some(true),
182 }
183 }
184
185 #[test]
186 fn new_rejects_empty_fields() {
187 let base = ClientCredentialsConfig {
188 token_url: "http://localhost/token".into(),
189 client_id: "id".into(),
190 client_secret: "secret".into(),
191 scopes: vec![],
192 refresh_skew_seconds: None,
193 accept_invalid_certs: Some(true),
194 };
195
196 let mut c = base.clone();
197 c.token_url = " ".into();
198 assert!(matches!(
199 OAuth2TokenProvider::new(c).unwrap_err(),
200 OAuth2TokenProviderError::InvalidConfig(_)
201 ));
202
203 let mut c = base.clone();
204 c.client_id = "".into();
205 assert!(matches!(
206 OAuth2TokenProvider::new(c).unwrap_err(),
207 OAuth2TokenProviderError::InvalidConfig(_)
208 ));
209
210 let mut c = base.clone();
211 c.client_secret = " ".into();
212 assert!(matches!(
213 OAuth2TokenProvider::new(c).unwrap_err(),
214 OAuth2TokenProviderError::InvalidConfig(_)
215 ));
216 }
217
218 #[test]
219 fn refresh_skew_none_uses_default() {
220 let config = ClientCredentialsConfig {
221 token_url: "http://localhost/token".into(),
222 client_id: "id".into(),
223 client_secret: "secret".into(),
224 scopes: vec![],
225 refresh_skew_seconds: None,
226 accept_invalid_certs: Some(true),
227 };
228
229 let provider = OAuth2TokenProvider::new(config).unwrap();
230 assert_eq!(provider.refresh_skew(), DEFAULT_REFRESH_SKEW_SECONDS);
231 }
232
233 #[test]
234 fn get_token_fetches_with_request_body_auth_and_caches_token() {
235 let server = MockServer::start();
236
237 let token_mock = server.mock(|when, then| {
238 when.method(POST)
239 .path("/token")
240 .header("content-type", "application/x-www-form-urlencoded")
241 .body("grant_type=client_credentials&scope=scope-a+scope-b&client_id=client-id&client_secret=client-secret");
242
243 then.status(200)
244 .header("content-type", "application/json")
245 .body(r#"{"access_token":"t1","token_type":"bearer","expires_in":3600}"#);
246 });
247
248 let clock = Arc::new(ManualClock::new(1000));
249 let provider =
250 OAuth2TokenProvider::with_clock(cfg(format!("{}/token", server.base_url())), clock)
251 .unwrap();
252
253 let t1 = provider.get_token().unwrap();
254 let t2 = provider.get_token().unwrap();
255
256 assert_eq!(t1, "t1");
257 assert_eq!(t2, "t1");
258
259 token_mock.assert_calls(1);
260 }
261
262 #[test]
263 fn get_token_refreshes_when_skew_window_reached() {
264 let server = MockServer::start();
265
266 let mock_server = server.mock(|when, then| {
267 when.method(POST).path("/token");
268 then.status(200)
269 .header("content-type", "application/json")
270 .body(r#"{"access_token":"t1","token_type":"bearer","expires_in":10}"#);
271 });
272
273 let clock = Arc::new(ManualClock::new(1000));
274
275 let config = ClientCredentialsConfig {
276 token_url: format!("{}/token", server.base_url()),
277 client_id: "client-id".into(),
278 client_secret: "client-secret".into(),
279 scopes: vec![],
280 refresh_skew_seconds: Some(30),
281 accept_invalid_certs: Some(true),
282 };
283
284 let provider = OAuth2TokenProvider::with_clock(config, clock.clone()).unwrap();
285
286 let token = provider.get_token().unwrap();
287 assert_eq!(token, "t1");
288
289 clock.advance(300);
291
292 provider.get_token().unwrap();
293
294 mock_server.assert_calls(2);
295 }
296
297 #[test]
298 fn get_token_does_not_refresh_when_skew_is_zero() {
299 let server = MockServer::start();
300
301 let mock_server = server.mock(|when, then| {
302 when.method(POST).path("/token");
303 then.status(200)
304 .header("content-type", "application/json")
305 .body(r#"{"access_token":"t1","token_type":"bearer","expires_in":1}"#);
306 });
307
308 let clock = Arc::new(ManualClock::new(1000));
309
310 let config = ClientCredentialsConfig {
311 token_url: format!("{}/token", server.base_url()),
312 client_id: "client-id".into(),
313 client_secret: "client-secret".into(),
314 scopes: vec![],
315 refresh_skew_seconds: Some(0),
316 accept_invalid_certs: Some(true),
317 };
318
319 let provider = OAuth2TokenProvider::with_clock(config, clock.clone()).unwrap();
320
321 let token = provider.get_token().unwrap();
322 clock.advance(1000);
324 provider.get_token().unwrap();
325
326 assert_eq!(token, "t1");
327 mock_server.assert_calls(1);
328 }
329
330 #[test]
331 fn missing_expires_in_returns_error() {
332 let server = MockServer::start();
333
334 server.mock(|when, then| {
335 when.method(POST).path("/token");
336 then.status(200)
337 .header("content-type", "application/json")
338 .body(r#"{"access_token":"t1","token_type":"bearer"}"#);
339 });
340
341 let clock = Arc::new(ManualClock::new(1000));
342 let provider =
343 OAuth2TokenProvider::with_clock(cfg(format!("{}/token", server.base_url())), clock)
344 .unwrap();
345
346 let err = provider.get_token().unwrap_err();
347 assert!(matches!(err, OAuth2TokenProviderError::MissingExpiresIn));
348 }
349
350 #[test]
351 fn token_request_error_is_mapped() {
352 let clock = Arc::new(ManualClock::new(1000));
353 let provider = OAuth2TokenProvider::with_clock(
354 ClientCredentialsConfig {
355 token_url: "http://127.0.0.1:9/token".into(),
356 client_id: "client-id".into(),
357 client_secret: "client-secret".into(),
358 scopes: vec![],
359 refresh_skew_seconds: None,
360 accept_invalid_certs: Some(true),
361 },
362 clock,
363 )
364 .unwrap();
365
366 let err = provider.get_token().unwrap_err();
367 assert!(matches!(err, OAuth2TokenProviderError::TokenRequest(_)));
368 }
369}