Skip to main content

fraiseql_auth/
multi_provider.rs

1//! Multi-provider authentication — unified entry point for social login.
2//!
3//! Enables `GET /auth/v1/authorize?provider=github&redirect_uri=...` with
4//! automatic provider resolution and state-encoded provider tracking through
5//! the OAuth callback.
6
7use 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
22/// Maximum length for the `redirect_uri` query parameter.
23const MAX_REDIRECT_URI_BYTES: usize = 2_048;
24
25/// Maximum length for the `provider` query parameter.
26const MAX_PROVIDER_NAME_BYTES: usize = 128;
27
28/// Shared state for the multi-provider auth endpoints.
29#[derive(Clone)]
30pub struct MultiProviderAuthState {
31    /// OAuth providers keyed by name (e.g., "github", "google").
32    providers:     HashMap<String, Arc<dyn OAuthProvider>>,
33    /// CSRF state store (in-memory or Redis).
34    state_store:   Arc<dyn StateStore>,
35    /// Session backend for creating sessions after successful auth.
36    session_store: Arc<dyn SessionStore>,
37    /// Optional user store for account linking (same email → same user).
38    user_store:    Option<Arc<dyn AccountStore>>,
39}
40
41impl MultiProviderAuthState {
42    /// Create a new multi-provider auth state.
43    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    /// Set the user store for account linking.
53    ///
54    /// When set, the callback handler uses [`AccountStore::link_or_create_user`] to
55    /// resolve provider identities to local users, enabling automatic account
56    /// linking when the same email appears across different providers.
57    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    /// Register an OAuth provider under the given name.
63    pub fn register_provider(&mut self, name: impl Into<String>, provider: Arc<dyn OAuthProvider>) {
64        self.providers.insert(name.into(), provider);
65    }
66
67    /// List the names of all registered providers.
68    #[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    /// Look up a provider by name.
76    #[must_use]
77    pub fn get_provider(&self, name: &str) -> Option<&Arc<dyn OAuthProvider>> {
78        self.providers.get(name)
79    }
80}
81
82// ---------------------------------------------------------------------------
83// Query / response types
84// ---------------------------------------------------------------------------
85
86/// Query parameters for `GET /auth/v1/authorize`.
87#[derive(Debug, Deserialize)]
88pub struct AuthorizeQuery {
89    /// Provider name (e.g., "github", "google").
90    pub provider:     String,
91    /// Client application callback URI.
92    pub redirect_uri: String,
93}
94
95/// Query parameters for `GET /auth/v1/callback`.
96#[derive(Debug, Deserialize)]
97pub struct CallbackQuery {
98    /// Authorization code from the provider.
99    pub code:              Option<String>,
100    /// CSRF state token.
101    pub state:             Option<String>,
102    /// Provider error code.
103    pub error:             Option<String>,
104    /// Provider error description.
105    pub error_description: Option<String>,
106}
107
108/// Response for `GET /auth/v1/providers`.
109#[derive(Debug, Serialize)]
110pub struct ProvidersResponse {
111    /// Available provider names.
112    pub providers: Vec<String>,
113}
114
115/// Token response returned after a successful callback.
116#[derive(Debug, Serialize)]
117pub struct AuthTokenResponse {
118    /// Access token for API requests.
119    pub access_token:  String,
120    /// Refresh token (if available).
121    #[serde(skip_serializing_if = "Option::is_none")]
122    pub refresh_token: Option<String>,
123    /// Token type (always "Bearer").
124    pub token_type:    String,
125    /// Seconds until the access token expires.
126    pub expires_in:    u64,
127    /// Provider that authenticated the user.
128    pub provider:      String,
129}
130
131impl AuthTokenResponse {
132    /// Returns a builder for `AuthTokenResponse`.
133    #[must_use = "builder does nothing until .build() is called"]
134    pub fn builder() -> AuthTokenResponseBuilder {
135        AuthTokenResponseBuilder::default()
136    }
137}
138
139/// Builder for [`AuthTokenResponse`].
140#[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    /// Sets the access token.
151    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    /// Sets the refresh token.
157    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    /// Sets the token type (typically `"Bearer"`).
163    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    /// Sets the number of seconds until the access token expires.
169    #[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    /// Sets the provider that authenticated the user.
176    pub fn provider(mut self, provider: impl Into<String>) -> Self {
177        self.provider = Some(provider.into());
178        self
179    }
180
181    /// Builds the [`AuthTokenResponse`].
182    ///
183    /// # Errors
184    ///
185    /// Returns an error string if any required field (`access_token`, `token_type`,
186    /// `expires_in`, or `provider`) was not set.
187    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
200// ---------------------------------------------------------------------------
201// Helpers
202// ---------------------------------------------------------------------------
203
204fn json_error(status: StatusCode, message: &str) -> Response {
205    (status, Json(serde_json::json!({ "error": message }))).into_response()
206}
207
208// ---------------------------------------------------------------------------
209// GET /auth/v1/providers
210// ---------------------------------------------------------------------------
211
212/// List available authentication providers.
213///
214/// # Responses
215///
216/// - `200` JSON `{ providers: ["github", "google", ...] }`
217pub async fn list_providers(
218    State(state): State<Arc<MultiProviderAuthState>>,
219) -> Json<ProvidersResponse> {
220    Json(ProvidersResponse {
221        providers: state.provider_names(),
222    })
223}
224
225// ---------------------------------------------------------------------------
226// GET /auth/v1/authorize
227// ---------------------------------------------------------------------------
228
229/// Initiate the OAuth flow for a specific provider.
230///
231/// Generates a CSRF state token, stores it with the provider name, then
232/// redirects to the provider's authorization URL.
233///
234/// # Query parameters
235///
236/// - `provider` — **required**: provider name (must match a registered provider).
237/// - `redirect_uri` — **required**: client application callback URI. It is validated for presence
238///   and length but is **not currently used for a server-side redirect**: [`callback`] returns the
239///   session tokens as JSON for the client to handle. A server-side redirect to this URI is
240///   intentionally not implemented yet because it would be an open-redirect vector without a
241///   configured allow-list of permitted redirect URIs. Tracked as a follow-up feature in #427
242///   (allow-list-backed redirect flow).
243///
244/// # Responses
245///
246/// - `302` — redirect to the provider's authorization endpoint.
247/// - `400` — missing or invalid parameters, unknown provider.
248///
249/// # Errors
250///
251/// Returns a `400` JSON error if the provider is unknown, redirect_uri is empty/oversized,
252/// or the state store is at capacity.
253pub async fn authorize(
254    State(state): State<Arc<MultiProviderAuthState>>,
255    Query(q): Query<AuthorizeQuery>,
256) -> Response {
257    // Validate provider name length
258    if q.provider.len() > MAX_PROVIDER_NAME_BYTES {
259        return json_error(StatusCode::BAD_REQUEST, "provider name exceeds maximum length");
260    }
261
262    // Validate redirect_uri
263    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    // Look up provider
271    let Some(provider) = state.get_provider(&q.provider) else {
272        return json_error(StatusCode::BAD_REQUEST, &format!("unknown provider: {}", q.provider));
273    };
274
275    // Generate state and store with provider name
276    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; // 10 minutes
286
287    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    // Generate authorization URL
296    let authorization_url = provider.authorization_url(&state_value);
297
298    Redirect::to(&authorization_url).into_response()
299}
300
301// ---------------------------------------------------------------------------
302// GET /auth/v1/callback
303// ---------------------------------------------------------------------------
304
305/// Complete the OAuth flow after the provider redirects back.
306///
307/// Validates the state token, resolves the provider from the stored state,
308/// exchanges the authorization code for tokens, retrieves user info, and
309/// creates a session.
310///
311/// # Query parameters
312///
313/// - `code` — authorization code from the provider.
314/// - `state` — CSRF state token.
315///
316/// # Responses
317///
318/// - `200` JSON `{ access_token, refresh_token?, token_type, expires_in, provider }`
319/// - `400` — invalid state, missing parameters, or provider error.
320/// - `502` — token exchange with the provider failed.
321///
322/// # Errors
323///
324/// Returns `400` if the state is invalid/expired, code is missing, or the provider
325/// returned an error. Returns `502` if the token exchange or user info fetch fails.
326#[allow(clippy::cognitive_complexity)] // Reason: OAuth callback with state validation, token exchange, user info, and session creation
327pub async fn callback(
328    State(state): State<Arc<MultiProviderAuthState>>,
329    Query(q): Query<CallbackQuery>,
330) -> Response {
331    // Surface provider errors
332    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    // Validate required parameters
346    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    // Consume state (atomic remove) and get provider name
351    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    // Check state expiry. Fail-closed: if the clock cannot be read, reject rather than
356    // treat the (possibly expired) CSRF state as valid (matches the authorize path).
357    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    // Look up provider
369    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    // Exchange code for tokens
375    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    // Get user info from provider
384    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    // Resolve local user ID — use AccountStore for account linking when available,
393    // otherwise fall back to raw provider user ID.
394    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    // Create session (7-day expiry)
415    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// ---------------------------------------------------------------------------
439// Tests
440// ---------------------------------------------------------------------------