1use std::{collections::HashMap, sync::Arc};
8
9use axum::{
10 Json,
11 extract::{Query, State},
12 http::StatusCode,
13 response::{IntoResponse, Redirect, Response},
14};
15use serde::{Deserialize, Serialize};
16
17use crate::{
18 account_linking::AccountStore, handlers::generate_secure_state, provider::OAuthProvider,
19 session::SessionStore, state_store::StateStore,
20};
21
22const MAX_REDIRECT_URI_BYTES: usize = 2_048;
24
25const MAX_PROVIDER_NAME_BYTES: usize = 128;
27
28#[derive(Clone)]
30pub struct MultiProviderAuthState {
31 providers: HashMap<String, Arc<dyn OAuthProvider>>,
33 state_store: Arc<dyn StateStore>,
35 session_store: Arc<dyn SessionStore>,
37 user_store: Option<Arc<dyn AccountStore>>,
39}
40
41impl MultiProviderAuthState {
42 pub fn new(state_store: Arc<dyn StateStore>, session_store: Arc<dyn SessionStore>) -> Self {
44 Self {
45 providers: HashMap::new(),
46 state_store,
47 session_store,
48 user_store: None,
49 }
50 }
51
52 pub fn with_user_store(mut self, user_store: Arc<dyn AccountStore>) -> Self {
58 self.user_store = Some(user_store);
59 self
60 }
61
62 pub fn register_provider(&mut self, name: impl Into<String>, provider: Arc<dyn OAuthProvider>) {
64 self.providers.insert(name.into(), provider);
65 }
66
67 #[must_use]
69 pub fn provider_names(&self) -> Vec<String> {
70 let mut names: Vec<String> = self.providers.keys().cloned().collect();
71 names.sort();
72 names
73 }
74
75 #[must_use]
77 pub fn get_provider(&self, name: &str) -> Option<&Arc<dyn OAuthProvider>> {
78 self.providers.get(name)
79 }
80}
81
82#[derive(Debug, Deserialize)]
88pub struct AuthorizeQuery {
89 pub provider: String,
91 pub redirect_uri: String,
93}
94
95#[derive(Debug, Deserialize)]
97pub struct CallbackQuery {
98 pub code: Option<String>,
100 pub state: Option<String>,
102 pub error: Option<String>,
104 pub error_description: Option<String>,
106}
107
108#[derive(Debug, Serialize)]
110pub struct ProvidersResponse {
111 pub providers: Vec<String>,
113}
114
115#[derive(Debug, Serialize)]
117pub struct AuthTokenResponse {
118 pub access_token: String,
120 #[serde(skip_serializing_if = "Option::is_none")]
122 pub refresh_token: Option<String>,
123 pub token_type: String,
125 pub expires_in: u64,
127 pub provider: String,
129}
130
131impl AuthTokenResponse {
132 #[must_use = "builder does nothing until .build() is called"]
134 pub fn builder() -> AuthTokenResponseBuilder {
135 AuthTokenResponseBuilder::default()
136 }
137}
138
139#[derive(Debug, Default)]
141pub struct AuthTokenResponseBuilder {
142 access_token: Option<String>,
143 refresh_token: Option<String>,
144 token_type: Option<String>,
145 expires_in: Option<u64>,
146 provider: Option<String>,
147}
148
149impl AuthTokenResponseBuilder {
150 pub fn access_token(mut self, access_token: impl Into<String>) -> Self {
152 self.access_token = Some(access_token.into());
153 self
154 }
155
156 pub fn refresh_token(mut self, refresh_token: impl Into<String>) -> Self {
158 self.refresh_token = Some(refresh_token.into());
159 self
160 }
161
162 pub fn token_type(mut self, token_type: impl Into<String>) -> Self {
164 self.token_type = Some(token_type.into());
165 self
166 }
167
168 #[must_use = "builder method returns modified builder"]
170 pub const fn expires_in(mut self, expires_in: u64) -> Self {
171 self.expires_in = Some(expires_in);
172 self
173 }
174
175 pub fn provider(mut self, provider: impl Into<String>) -> Self {
177 self.provider = Some(provider.into());
178 self
179 }
180
181 pub fn build(self) -> Result<AuthTokenResponse, String> {
188 Ok(AuthTokenResponse {
189 access_token: self
190 .access_token
191 .ok_or("AuthTokenResponse: access_token is required")?,
192 refresh_token: self.refresh_token,
193 token_type: self.token_type.ok_or("AuthTokenResponse: token_type is required")?,
194 expires_in: self.expires_in.ok_or("AuthTokenResponse: expires_in is required")?,
195 provider: self.provider.ok_or("AuthTokenResponse: provider is required")?,
196 })
197 }
198}
199
200fn json_error(status: StatusCode, message: &str) -> Response {
205 (status, Json(serde_json::json!({ "error": message }))).into_response()
206}
207
208pub async fn list_providers(
218 State(state): State<Arc<MultiProviderAuthState>>,
219) -> Json<ProvidersResponse> {
220 Json(ProvidersResponse {
221 providers: state.provider_names(),
222 })
223}
224
225pub async fn authorize(
254 State(state): State<Arc<MultiProviderAuthState>>,
255 Query(q): Query<AuthorizeQuery>,
256) -> Response {
257 if q.provider.len() > MAX_PROVIDER_NAME_BYTES {
259 return json_error(StatusCode::BAD_REQUEST, "provider name exceeds maximum length");
260 }
261
262 if q.redirect_uri.is_empty() {
264 return json_error(StatusCode::BAD_REQUEST, "redirect_uri is required");
265 }
266 if q.redirect_uri.len() > MAX_REDIRECT_URI_BYTES {
267 return json_error(StatusCode::BAD_REQUEST, "redirect_uri exceeds maximum length");
268 }
269
270 let Some(provider) = state.get_provider(&q.provider) else {
272 return json_error(StatusCode::BAD_REQUEST, &format!("unknown provider: {}", q.provider));
273 };
274
275 let state_value = generate_secure_state();
277
278 let Ok(now) = std::time::SystemTime::now()
279 .duration_since(std::time::UNIX_EPOCH)
280 .map(|d| d.as_secs())
281 else {
282 return json_error(StatusCode::INTERNAL_SERVER_ERROR, "system clock error");
283 };
284
285 let expiry = now + 600; if let Err(e) = state.state_store.store(state_value.clone(), q.provider.clone(), expiry).await {
288 tracing::error!("state store failed: {e}");
289 return json_error(
290 StatusCode::INTERNAL_SERVER_ERROR,
291 "authorization flow could not be started",
292 );
293 }
294
295 let authorization_url = provider.authorization_url(&state_value);
297
298 Redirect::to(&authorization_url).into_response()
299}
300
301#[allow(clippy::cognitive_complexity)] pub async fn callback(
328 State(state): State<Arc<MultiProviderAuthState>>,
329 Query(q): Query<CallbackQuery>,
330) -> Response {
331 if let Some(err) = q.error {
333 let desc = q.error_description.as_deref().unwrap_or("(no description)");
334 tracing::warn!(provider_error = %err, description = %desc, "OAuth provider returned error");
335 let client_message = match err.as_str() {
336 "access_denied" => "Access was denied",
337 "login_required" => "Authentication is required",
338 "invalid_request" | "invalid_scope" => "Invalid authorization request",
339 "server_error" | "temporarily_unavailable" => "Authorization server error",
340 _ => "Authorization failed",
341 };
342 return json_error(StatusCode::BAD_REQUEST, client_message);
343 }
344
345 let (Some(code), Some(state_token)) = (q.code, q.state) else {
347 return json_error(StatusCode::BAD_REQUEST, "missing code or state parameter");
348 };
349
350 let Ok((provider_name, expiry)) = state.state_store.retrieve(&state_token).await else {
352 return json_error(StatusCode::BAD_REQUEST, "invalid or expired state token");
353 };
354
355 let Ok(now) = std::time::SystemTime::now()
358 .duration_since(std::time::UNIX_EPOCH)
359 .map(|d| d.as_secs())
360 else {
361 return json_error(StatusCode::INTERNAL_SERVER_ERROR, "system clock error");
362 };
363
364 if now > expiry {
365 return json_error(StatusCode::BAD_REQUEST, "state token expired");
366 }
367
368 let Some(provider) = state.get_provider(&provider_name) else {
370 tracing::error!(provider = %provider_name, "provider from state not found in registry");
371 return json_error(StatusCode::INTERNAL_SERVER_ERROR, "provider configuration error");
372 };
373
374 let token_response = match provider.exchange_code(&code).await {
376 Ok(t) => t,
377 Err(e) => {
378 tracing::error!(error = %e, "token exchange failed");
379 return json_error(StatusCode::BAD_GATEWAY, "token exchange with provider failed");
380 },
381 };
382
383 let user_info = match provider.user_info(&token_response.access_token).await {
385 Ok(u) => u,
386 Err(e) => {
387 tracing::error!(error = %e, "user info fetch failed");
388 return json_error(StatusCode::BAD_GATEWAY, "failed to retrieve user information");
389 },
390 };
391
392 let local_user_id = if let Some(account_store) = &state.user_store {
395 match account_store
396 .link_or_create_user(
397 user_info.email.as_deref(),
398 user_info.email_verified,
399 &provider_name,
400 &user_info.id,
401 )
402 .await
403 {
404 Ok(result) => result.user_id,
405 Err(e) => {
406 tracing::error!(error = %e, "account store lookup failed");
407 return json_error(StatusCode::INTERNAL_SERVER_ERROR, "user resolution failed");
408 },
409 }
410 } else {
411 user_info.id.clone()
412 };
413
414 let session_expiry = now + (7 * 24 * 60 * 60);
416 let session_tokens = match state
417 .session_store
418 .create_session(&local_user_id, session_expiry)
419 .await
420 {
421 Ok(t) => t,
422 Err(e) => {
423 tracing::error!(error = %e, "session creation failed");
424 return json_error(StatusCode::INTERNAL_SERVER_ERROR, "session could not be created");
425 },
426 };
427
428 Json(AuthTokenResponse {
429 access_token: session_tokens.access_token,
430 refresh_token: Some(session_tokens.refresh_token),
431 token_type: "Bearer".to_string(),
432 expires_in: session_tokens.expires_in,
433 provider: provider_name,
434 })
435 .into_response()
436}
437
438