Skip to main content

gvm_auth/
oauth2.rs

1// SPDX-FileCopyrightText: 2026 Greenbone AG
2//
3// SPDX-License-Identifier: GPL-3.0-or-later
4
5use 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)) // only for dev, same as curl -k
131            .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        // Move time forward a bit
290        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        // Move time forward a bit
323        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}