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(
249 State(state): State<Arc<MultiProviderAuthState>>,
250 Query(q): Query<AuthorizeQuery>,
251) -> Response {
252 if q.provider.len() > MAX_PROVIDER_NAME_BYTES {
254 return json_error(StatusCode::BAD_REQUEST, "provider name exceeds maximum length");
255 }
256
257 if q.redirect_uri.is_empty() {
259 return json_error(StatusCode::BAD_REQUEST, "redirect_uri is required");
260 }
261 if q.redirect_uri.len() > MAX_REDIRECT_URI_BYTES {
262 return json_error(StatusCode::BAD_REQUEST, "redirect_uri exceeds maximum length");
263 }
264
265 let Some(provider) = state.get_provider(&q.provider) else {
267 return json_error(StatusCode::BAD_REQUEST, &format!("unknown provider: {}", q.provider));
268 };
269
270 let state_value = generate_secure_state();
272
273 let Ok(now) = std::time::SystemTime::now()
274 .duration_since(std::time::UNIX_EPOCH)
275 .map(|d| d.as_secs())
276 else {
277 return json_error(StatusCode::INTERNAL_SERVER_ERROR, "system clock error");
278 };
279
280 let expiry = now + 600; if let Err(e) = state.state_store.store(state_value.clone(), q.provider.clone(), expiry).await {
283 tracing::error!("state store failed: {e}");
284 return json_error(
285 StatusCode::INTERNAL_SERVER_ERROR,
286 "authorization flow could not be started",
287 );
288 }
289
290 let authorization_url = provider.authorization_url(&state_value);
292
293 Redirect::to(&authorization_url).into_response()
294}
295
296#[allow(clippy::cognitive_complexity)] pub async fn callback(
323 State(state): State<Arc<MultiProviderAuthState>>,
324 Query(q): Query<CallbackQuery>,
325) -> Response {
326 if let Some(err) = q.error {
328 let desc = q.error_description.as_deref().unwrap_or("(no description)");
329 tracing::warn!(provider_error = %err, description = %desc, "OAuth provider returned error");
330 let client_message = match err.as_str() {
331 "access_denied" => "Access was denied",
332 "login_required" => "Authentication is required",
333 "invalid_request" | "invalid_scope" => "Invalid authorization request",
334 "server_error" | "temporarily_unavailable" => "Authorization server error",
335 _ => "Authorization failed",
336 };
337 return json_error(StatusCode::BAD_REQUEST, client_message);
338 }
339
340 let (Some(code), Some(state_token)) = (q.code, q.state) else {
342 return json_error(StatusCode::BAD_REQUEST, "missing code or state parameter");
343 };
344
345 let Ok((provider_name, expiry)) = state.state_store.retrieve(&state_token).await else {
347 return json_error(StatusCode::BAD_REQUEST, "invalid or expired state token");
348 };
349
350 let now = std::time::SystemTime::now()
352 .duration_since(std::time::UNIX_EPOCH)
353 .unwrap_or_default()
354 .as_secs();
355
356 if now > expiry {
357 return json_error(StatusCode::BAD_REQUEST, "state token expired");
358 }
359
360 let Some(provider) = state.get_provider(&provider_name) else {
362 tracing::error!(provider = %provider_name, "provider from state not found in registry");
363 return json_error(StatusCode::INTERNAL_SERVER_ERROR, "provider configuration error");
364 };
365
366 let token_response = match provider.exchange_code(&code).await {
368 Ok(t) => t,
369 Err(e) => {
370 tracing::error!(error = %e, "token exchange failed");
371 return json_error(StatusCode::BAD_GATEWAY, "token exchange with provider failed");
372 },
373 };
374
375 let user_info = match provider.user_info(&token_response.access_token).await {
377 Ok(u) => u,
378 Err(e) => {
379 tracing::error!(error = %e, "user info fetch failed");
380 return json_error(StatusCode::BAD_GATEWAY, "failed to retrieve user information");
381 },
382 };
383
384 let local_user_id = if let Some(account_store) = &state.user_store {
387 match account_store
388 .link_or_create_user(&user_info.email, &provider_name, &user_info.id)
389 .await
390 {
391 Ok(result) => result.user_id,
392 Err(e) => {
393 tracing::error!(error = %e, "account store lookup failed");
394 return json_error(StatusCode::INTERNAL_SERVER_ERROR, "user resolution failed");
395 },
396 }
397 } else {
398 user_info.id.clone()
399 };
400
401 let session_expiry = now + (7 * 24 * 60 * 60);
403 let session_tokens = match state
404 .session_store
405 .create_session(&local_user_id, session_expiry)
406 .await
407 {
408 Ok(t) => t,
409 Err(e) => {
410 tracing::error!(error = %e, "session creation failed");
411 return json_error(StatusCode::INTERNAL_SERVER_ERROR, "session could not be created");
412 },
413 };
414
415 Json(AuthTokenResponse {
416 access_token: session_tokens.access_token,
417 refresh_token: Some(session_tokens.refresh_token),
418 token_type: "Bearer".to_string(),
419 expires_in: session_tokens.expires_in,
420 provider: provider_name,
421 })
422 .into_response()
423}
424
425