1use std::{fmt, net::SocketAddr, sync::Arc};
4
5use axum::{
6 Json,
7 extract::{ConnectInfo, Query, State},
8 http::StatusCode,
9 response::IntoResponse,
10};
11use serde::{Deserialize, Serialize};
12
13use crate::{
14 audit::logger::{AuditEventType, SecretType, get_audit_logger},
15 error::{AuthError, Result},
16 provider::OAuthProvider,
17 rate_limiting::RateLimiters,
18 session::SessionStore,
19 state_store::StateStore,
20};
21
22#[derive(Clone)]
24pub struct AuthState {
25 pub oauth_provider: Arc<dyn OAuthProvider>,
27 pub session_store: Arc<dyn SessionStore>,
29 pub state_store: Arc<dyn StateStore>,
31 pub rate_limiters: Arc<RateLimiters>,
33}
34
35#[derive(Debug, Deserialize)]
37pub struct AuthStartRequest {
38 pub provider: Option<String>,
40}
41
42#[derive(Debug, Serialize)]
44pub struct AuthStartResponse {
45 pub authorization_url: String,
47}
48
49#[derive(Debug, Deserialize)]
51pub struct AuthCallbackQuery {
52 pub code: String,
54 pub state: String,
56 pub error: Option<String>,
58 pub error_description: Option<String>,
60}
61
62#[derive(Serialize)]
75#[non_exhaustive]
76#[doc(hidden)] pub struct AuthCallbackResponse {
78 pub access_token: String,
80 pub refresh_token: Option<String>,
82 pub token_type: String,
84 pub expires_in: u64,
86}
87
88impl fmt::Debug for AuthCallbackResponse {
89 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
90 f.debug_struct("AuthCallbackResponse")
94 .field("access_token", &"<redacted>")
95 .field("refresh_token", &self.refresh_token.as_ref().map(|_| "<redacted>"))
96 .field("token_type", &self.token_type)
97 .field("expires_in", &self.expires_in)
98 .finish()
99 }
100}
101
102impl AuthCallbackResponse {
103 #[must_use]
105 pub const fn new(
106 access_token: String,
107 refresh_token: Option<String>,
108 token_type: String,
109 expires_in: u64,
110 ) -> Self {
111 Self {
112 access_token,
113 refresh_token,
114 token_type,
115 expires_in,
116 }
117 }
118}
119
120#[derive(Debug, Deserialize)]
122pub struct AuthRefreshRequest {
123 pub refresh_token: String,
125}
126
127#[derive(Serialize)]
134pub struct AuthRefreshResponse {
135 pub access_token: String,
137 pub token_type: String,
139 pub expires_in: u64,
141}
142
143impl fmt::Debug for AuthRefreshResponse {
144 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
145 f.debug_struct("AuthRefreshResponse")
148 .field("access_token", &"<redacted>")
149 .field("token_type", &self.token_type)
150 .field("expires_in", &self.expires_in)
151 .finish()
152 }
153}
154
155#[derive(Debug, Deserialize)]
157pub struct AuthLogoutRequest {
158 pub refresh_token: Option<String>,
160}
161
162pub async fn auth_start(
178 State(state): State<AuthState>,
179 ConnectInfo(addr): ConnectInfo<SocketAddr>,
180 Json(req): Json<AuthStartRequest>,
181) -> Result<Json<AuthStartResponse>> {
182 let client_ip = addr.ip().to_string();
184 if state.rate_limiters.auth_start.check(&client_ip).is_err() {
185 return Err(AuthError::RateLimited {
186 retry_after_secs: state.rate_limiters.auth_start.clone_config().window_secs,
187 });
188 }
189
190 let state_value = generate_secure_state();
192
193 let now = std::time::SystemTime::now()
195 .duration_since(std::time::UNIX_EPOCH)
196 .map_err(|_| AuthError::SystemTimeError {
197 message: "Failed to get current system time".to_string(),
198 })?
199 .as_secs();
200
201 let expiry = now + 600;
203
204 let provider = req.provider.unwrap_or_else(|| "default".to_string());
206 state.state_store.store(state_value.clone(), provider, expiry).await?;
207
208 let authorization_url = state.oauth_provider.authorization_url(&state_value);
210
211 Ok(Json(AuthStartResponse { authorization_url }))
212}
213
214pub async fn auth_callback(
231 State(state): State<AuthState>,
232 ConnectInfo(addr): ConnectInfo<SocketAddr>,
233 Query(query): Query<AuthCallbackQuery>,
234) -> Result<impl IntoResponse> {
235 let client_ip = addr.ip().to_string();
237 if state.rate_limiters.auth_callback.check(&client_ip).is_err() {
238 return Err(AuthError::RateLimited {
239 retry_after_secs: state.rate_limiters.auth_callback.clone_config().window_secs,
240 });
241 }
242
243 validate_auth_input_len(&query.code, MAX_AUTH_CODE_BYTES, "code")?;
245 validate_auth_input_len(&query.state, MAX_STATE_BYTES, "state")?;
246
247 if let Some(error) = query.error {
249 let audit_logger = get_audit_logger();
250 audit_logger.log_failure(
251 AuditEventType::OauthCallback,
252 SecretType::AuthorizationCode,
253 None,
254 "exchange",
255 &error,
256 );
257 return Err(AuthError::OAuthError {
258 message: format!("{}: {}", error, query.error_description.unwrap_or_default()),
259 });
260 }
261
262 let (_provider_name, expiry) = state.state_store.retrieve(&query.state).await?;
264
265 let now = std::time::SystemTime::now()
267 .duration_since(std::time::UNIX_EPOCH)
268 .map_err(|_| AuthError::SystemTimeError {
269 message: "Failed to get current system time".to_string(),
270 })?
271 .as_secs();
272
273 if now > expiry {
274 let audit_logger = get_audit_logger();
275 audit_logger.log_failure(
276 AuditEventType::CsrfStateValidated,
277 SecretType::StateToken,
278 None,
279 "validate",
280 "State token expired",
281 );
282 return Err(AuthError::InvalidState);
283 }
284
285 let audit_logger = get_audit_logger();
287 audit_logger.log_success(
288 AuditEventType::CsrfStateValidated,
289 SecretType::StateToken,
290 None,
291 "validate",
292 );
293
294 let token_response = state.oauth_provider.exchange_code(&query.code).await?;
296
297 let audit_logger = get_audit_logger();
299 audit_logger.log_success(
300 AuditEventType::OauthCallback,
301 SecretType::AuthorizationCode,
302 None,
303 "exchange",
304 );
305
306 let user_info = state.oauth_provider.user_info(&token_response.access_token).await?;
308
309 let expires_at = now + (7 * 24 * 60 * 60);
311 let session_tokens = state.session_store.create_session(&user_info.id, expires_at).await?;
312
313 let audit_logger = get_audit_logger();
315 audit_logger.log_success(
316 AuditEventType::SessionTokenCreated,
317 SecretType::SessionToken,
318 Some(user_info.id.clone()),
319 "create",
320 );
321
322 let audit_logger = get_audit_logger();
324 audit_logger.log_success(
325 AuditEventType::AuthSuccess,
326 SecretType::SessionToken,
327 Some(user_info.id),
328 "oauth_flow",
329 );
330
331 let response = AuthCallbackResponse {
332 access_token: session_tokens.access_token,
333 refresh_token: Some(session_tokens.refresh_token),
334 token_type: "Bearer".to_string(),
335 expires_in: session_tokens.expires_in,
336 };
337
338 Ok(Json(response))
341}
342
343pub async fn auth_refresh(
359 State(state): State<AuthState>,
360 Json(req): Json<AuthRefreshRequest>,
361) -> Result<Json<AuthRefreshResponse>> {
362 use crate::session::hash_token;
363
364 validate_auth_input_len(&req.refresh_token, MAX_REFRESH_TOKEN_BYTES, "refresh_token")?;
366
367 let token_hash = hash_token(&req.refresh_token);
369 let session = state.session_store.get_session(&token_hash).await?;
370
371 if session.is_expired() {
375 let audit_logger = get_audit_logger();
376 audit_logger.log_failure(
377 AuditEventType::JwtRefresh,
378 SecretType::RefreshToken,
379 Some(session.user_id),
380 "refresh",
381 "Session expired",
382 );
383 return Err(AuthError::TokenExpired);
384 }
385
386 if state.rate_limiters.auth_refresh.check(&session.user_id).is_err() {
388 return Err(AuthError::RateLimited {
389 retry_after_secs: state.rate_limiters.auth_refresh.clone_config().window_secs,
390 });
391 }
392
393 let audit_logger = get_audit_logger();
395 audit_logger.log_success(
396 AuditEventType::SessionTokenValidation,
397 SecretType::RefreshToken,
398 Some(session.user_id),
399 "validate",
400 );
401
402 Err(AuthError::Internal {
405 message: "JWT signing not yet implemented — configure an OIDC provider for token issuance"
406 .to_string(),
407 })
408}
409
410pub async fn auth_logout(
425 State(state): State<AuthState>,
426 ConnectInfo(addr): ConnectInfo<SocketAddr>,
427 Json(req): Json<AuthLogoutRequest>,
428) -> Result<StatusCode> {
429 let client_ip = addr.ip().to_string();
430
431 if let Some(refresh_token) = req.refresh_token {
432 use crate::session::hash_token;
433 let token_hash = hash_token(&refresh_token);
434
435 let session = state.session_store.get_session(&token_hash).await?;
437
438 if state.rate_limiters.auth_logout.check(&session.user_id).is_err() {
440 return Err(AuthError::RateLimited {
441 retry_after_secs: state.rate_limiters.auth_logout.clone_config().window_secs,
442 });
443 }
444
445 state.session_store.revoke_session(&token_hash).await?;
446
447 let audit_logger = get_audit_logger();
449 audit_logger.log_success(
450 AuditEventType::SessionTokenRevoked,
451 SecretType::RefreshToken,
452 Some(session.user_id),
453 "revoke",
454 );
455 } else {
456 if state.rate_limiters.auth_logout.check(&client_ip).is_err() {
458 return Err(AuthError::RateLimited {
459 retry_after_secs: state.rate_limiters.auth_logout.clone_config().window_secs,
460 });
461 }
462 }
463
464 Ok(StatusCode::NO_CONTENT)
465}
466
467#[must_use]
470pub fn generate_secure_state() -> String {
471 use rand::RngCore as _;
472
473 let mut bytes = [0u8; 32];
475 rand::rng().fill_bytes(&mut bytes);
476
477 hex::encode(bytes)
479}
480
481pub const MAX_AUTH_CODE_BYTES: usize = 512;
488
489pub const MAX_STATE_BYTES: usize = 2_048;
497
498pub const MAX_REFRESH_TOKEN_BYTES: usize = 4_096;
503
504pub fn validate_auth_input_len(
514 input: &str,
515 max_bytes: usize,
516 field: &str,
517) -> crate::error::Result<()> {
518 if input.len() > max_bytes {
519 return Err(crate::error::AuthError::InvalidToken {
520 reason: format!("{field} exceeds maximum length ({} > {max_bytes} bytes)", input.len()),
521 });
522 }
523 Ok(())
524}
525
526#[cfg(test)]
527mod debug_redaction_tests {
528 use super::*;
532
533 const SECRET_ACCESS: &str = "eyJhbGciOiJIUzI1NiJ9.SUPER-SECRET-ACCESS-TOKEN.sig";
534 const SECRET_REFRESH: &str = "RT-SUPER-SECRET-REFRESH-TOKEN-do-not-leak";
535
536 #[test]
537 fn auth_callback_response_debug_redacts_access_and_refresh_tokens() {
538 let resp = AuthCallbackResponse {
539 access_token: SECRET_ACCESS.to_string(),
540 refresh_token: Some(SECRET_REFRESH.to_string()),
541 token_type: "Bearer".to_string(),
542 expires_in: 3600,
543 };
544
545 let debug_output = format!("{resp:?}");
546
547 assert!(
548 !debug_output.contains(SECRET_ACCESS),
549 "access_token leaked in Debug output: {debug_output}",
550 );
551 assert!(
552 !debug_output.contains(SECRET_REFRESH),
553 "refresh_token leaked in Debug output: {debug_output}",
554 );
555 assert!(debug_output.contains("redacted"), "redaction marker missing: {debug_output}",);
556 assert!(debug_output.contains("Bearer"));
558 assert!(debug_output.contains("3600"));
559 }
560
561 #[test]
562 fn auth_callback_response_debug_with_no_refresh_token_shows_none() {
563 let resp = AuthCallbackResponse {
564 access_token: SECRET_ACCESS.to_string(),
565 refresh_token: None,
566 token_type: "Bearer".to_string(),
567 expires_in: 900,
568 };
569
570 let debug_output = format!("{resp:?}");
571
572 assert!(!debug_output.contains(SECRET_ACCESS));
573 assert!(
574 debug_output.contains("None"),
575 "expected None to appear when refresh_token absent, got: {debug_output}",
576 );
577 }
578
579 #[test]
580 fn auth_refresh_response_debug_redacts_access_token() {
581 let resp = AuthRefreshResponse {
582 access_token: SECRET_ACCESS.to_string(),
583 token_type: "Bearer".to_string(),
584 expires_in: 1800,
585 };
586
587 let debug_output = format!("{resp:?}");
588
589 assert!(
590 !debug_output.contains(SECRET_ACCESS),
591 "access_token leaked in Debug output: {debug_output}",
592 );
593 assert!(debug_output.contains("redacted"), "redaction marker missing: {debug_output}",);
594 assert!(debug_output.contains("Bearer"));
595 assert!(debug_output.contains("1800"));
596 }
597}