Skip to main content

headless_lms_server/domain/
authentication.rs

1//! Common functionality related to authenticating users.
2
3use crate::OAuthClient;
4use crate::config::server_runtime_config;
5use crate::domain::authorization::{AuthorizationToken, skip_authorize};
6use crate::prelude::*;
7use actix_http::Payload;
8use actix_session::Session;
9use actix_session::SessionExt;
10use actix_web::{FromRequest, HttpRequest};
11use anyhow::Result;
12use chrono::{DateTime, Duration, Utc};
13use futures::Future;
14use headless_lms_models::{self as models, users::User};
15use headless_lms_utils::http::REQWEST_CLIENT;
16use headless_lms_utils::services::tmc::TMCUser;
17use headless_lms_utils::services::tmc::TmcClient;
18use oauth2::EmptyExtraTokenFields;
19use oauth2::HttpClientError;
20use oauth2::RequestTokenError;
21use oauth2::ResourceOwnerPassword;
22use oauth2::ResourceOwnerUsername;
23use oauth2::StandardTokenResponse;
24use oauth2::TokenResponse;
25use oauth2::basic::BasicTokenType;
26use secrecy::ExposeSecret;
27use secrecy::SecretString;
28use serde::{Deserialize, Serialize};
29use serde_json::json;
30use sqlx::PgConnection;
31use std::pin::Pin;
32use subtle::ConstantTimeEq;
33use tracing_log::log;
34use uuid::Uuid;
35
36const SESSION_KEY: &str = "user";
37
38const MOOCFI_GRAPHQL_URL: &str = "https://www.mooc.fi/api";
39
40fn constant_time_eq_str(left: &str, right: &str) -> bool {
41    left.as_bytes().ct_eq(right.as_bytes()).into()
42}
43#[derive(Debug, Serialize, Deserialize)]
44struct GraphQLRequest<'a> {
45    query: &'a str,
46    #[serde(skip_serializing_if = "Option::is_none")]
47    variables: Option<serde_json::Value>,
48}
49
50#[derive(Debug, Serialize, Deserialize)]
51struct MoocfiUserResponse {
52    pub data: MoocfiUserResponseData,
53}
54
55#[derive(Debug, Serialize, Deserialize)]
56struct MoocfiUserResponseData {
57    pub user: MoocfiUserData,
58}
59
60#[derive(Debug, Serialize, Deserialize)]
61struct MoocfiUserData {
62    pub id: Uuid,
63}
64
65// upstream_id is private so FromRequest is the only way to construct an AuthUser.
66/// Extractor for an authenticated user.
67#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
68pub struct AuthUser {
69    pub id: Uuid,
70    pub created_at: DateTime<Utc>,
71    pub updated_at: DateTime<Utc>,
72    pub deleted_at: Option<DateTime<Utc>>,
73    pub fetched_from_db_at: Option<DateTime<Utc>>,
74    upstream_id: Option<i32>,
75}
76
77impl AuthUser {
78    /// The user's ID in TMC.
79    pub fn upstream_id(&self) -> Option<i32> {
80        self.upstream_id
81    }
82}
83
84/// The id of the user the session is signed in as, without checking the user still exists. For
85/// keying things like rate limits; use the [`AuthUser`] extractor to authenticate.
86pub fn session_user_id(session: &Session) -> Option<Uuid> {
87    session
88        .get::<AuthUser>(SESSION_KEY)
89        .ok()
90        .flatten()
91        .map(|user| user.id)
92}
93
94impl FromRequest for AuthUser {
95    type Error = ControllerError;
96    type Future = Pin<Box<dyn Future<Output = Result<Self, Self::Error>>>>;
97
98    fn from_request(req: &HttpRequest, _payload: &mut Payload) -> Self::Future {
99        let req = req.clone();
100        Box::pin(async move {
101            let req = req.clone();
102            let session = req.get_session();
103            let pool: Option<&web::Data<PgPool>> = req.app_data();
104            match session.get::<AuthUser>(SESSION_KEY) {
105                Ok(Some(user)) => Ok(verify_auth_user_exists(user, pool, &session).await?),
106                Ok(None) => Err(controller_err!(
107                    Unauthorized,
108                    "You are not currently logged in. Please sign in to continue.".to_string()
109                )),
110                Err(_) => {
111                    // session had an invalid value
112                    session.remove(SESSION_KEY);
113                    Err(controller_err!(
114                        Unauthorized,
115                        "Your session is invalid or has expired. Please sign in again.".to_string()
116                    ))
117                }
118            }
119        })
120    }
121}
122
123/**
124 * Re-fetches the user from the database and refreshes the session once it is more than 3 hours
125 * old; otherwise returns the session's cached AuthUser unchanged.
126 */
127async fn verify_auth_user_exists(
128    auth_user: AuthUser,
129    pool: Option<&web::Data<PgPool>>,
130    session: &Session,
131) -> Result<AuthUser, ControllerError> {
132    if let Some(fetched_from_db_at) = auth_user.fetched_from_db_at {
133        let time_now = Utc::now();
134        let time_hour_ago = time_now - Duration::hours(3);
135        if fetched_from_db_at > time_hour_ago {
136            return Ok(auth_user);
137        }
138    }
139    if let Some(pool) = pool {
140        info!("Checking whether the user saved in the session still exists in the database.");
141        let mut conn = pool.acquire().await?;
142        let user = models::users::get_by_id(&mut conn, auth_user.id).await?;
143        remember(session, user)?;
144        match session.get::<AuthUser>(SESSION_KEY) {
145            Ok(Some(session_user)) => Ok(session_user),
146            Ok(None) => Err(controller_err!(
147                InternalServerError,
148                "User did not persist in the session".to_string()
149            )),
150            Err(e) => Err(controller_err!(
151                InternalServerError,
152                "User did not persist in the session".to_string(),
153                e
154            )),
155        }
156    } else {
157        warn!("No database pool provided to verify_auth_user_exists");
158        Err(controller_err!(
159            InternalServerError,
160            "Unable to verify your user account. The database connection is unavailable."
161                .to_string()
162        ))
163    }
164}
165
166/// Stores the user as authenticated in the given session.
167pub fn remember(session: &Session, user: models::users::User) -> Result<()> {
168    let auth_user = AuthUser {
169        id: user.id,
170        created_at: user.created_at,
171        updated_at: user.updated_at,
172        deleted_at: user.deleted_at,
173        upstream_id: user.upstream_id,
174        fetched_from_db_at: Some(Utc::now()),
175    };
176    session
177        .insert(SESSION_KEY, auth_user)
178        .map_err(|_| anyhow::anyhow!("Failed to insert to session"))
179}
180
181/// Checks if the user is authenticated in the given session.
182pub async fn has_auth_user_session(session: &Session, pool: web::Data<PgPool>) -> bool {
183    match session.get::<AuthUser>(SESSION_KEY) {
184        Ok(Some(sesssion_auth_user)) => {
185            verify_auth_user_exists(sesssion_auth_user, Some(&pool), session)
186                .await
187                .is_ok()
188        }
189        _ => false,
190    }
191}
192
193/// Forgets authentication from the current session, if any.
194pub fn forget(session: &Session) {
195    session.purge();
196}
197
198/// Returns the bearer token only when there is no authenticated user, for the anonymous
199/// chatbot-embed path; a logged-in request's token is never surfaced here.
200pub fn handle_anonymous_token(req: &HttpRequest, user: Option<AuthUser>) -> Option<String> {
201    let anonymous_token_value = req
202        .headers()
203        .get("authorization")
204        .and_then(|anonymous_token| anonymous_token.to_str().ok()?.strip_prefix("Bearer "));
205
206    if let (Some(anonymous_token), None) = (anonymous_token_value, user) {
207        Some(anonymous_token.to_owned())
208    } else {
209        None
210    }
211}
212
213/// Checks the Authorization header against a secret from environment variables to verify the
214/// request originates from the TMC server.
215pub async fn authenticate_tmc_server(
216    request: &HttpRequest,
217) -> Result<AuthorizationToken, ControllerError> {
218    let tmc_server_secret_for_communicating_to_secret_project =
219        &server_runtime_config().tmc_server_secret_for_communicating_to_secret_project;
220    let auth_header = request
221        .headers()
222        .get("Authorization")
223        .ok_or_else(|| {
224            controller_err!(
225                Unauthorized,
226                "TMC server authorization failed: Missing Authorization header.".to_string()
227            )
228        })?
229        .to_str()
230        .map_err(|_| {
231            controller_err!(
232                Unauthorized,
233                "TMC server authorization failed: Invalid Authorization header format.".to_string()
234            )
235        })?;
236    if constant_time_eq_str(
237        auth_header,
238        tmc_server_secret_for_communicating_to_secret_project.expose_secret(),
239    ) {
240        return Ok(skip_authorize());
241    }
242    Err(controller_err!(
243        Unauthorized,
244        "TMC server authorization failed: Invalid authorization token.".to_string()
245    ))
246}
247
248pub fn parse_secret_key_from_header(header: &HttpRequest) -> Result<&str, ControllerError> {
249    let raw_token = header
250        .headers()
251        .get("Authorization")
252        .map_or(Ok(""), |x| x.to_str())
253        .map_err(|_| anyhow::anyhow!("Authorization header contains invalid characters."))?;
254    if !raw_token.starts_with("Basic") {
255        return Err(controller_err!(
256            Forbidden,
257            "Access denied: Authorization header must use Basic authentication format.".to_string()
258        ));
259    }
260    let secret_key = raw_token.split(' ').nth(1).ok_or_else(|| {
261        controller_err!(
262            Forbidden,
263            "Access denied: Malformed authorization token, expected 'Basic <token>' format."
264                .to_string()
265        )
266    })?;
267    Ok(secret_key)
268}
269
270/// Authenticates the user with mooc.fi, returning the authenticated user and their oauth token.
271pub async fn authenticate_tmc_mooc_fi_user(
272    conn: &mut PgConnection,
273    client: &OAuthClient,
274    email: String,
275    password: SecretString,
276    tmc_client: &TmcClient,
277) -> anyhow::Result<Option<(User, SecretString)>> {
278    info!("Attempting to authenticate user with TMC");
279    let token = match exchange_password_with_tmc(client, email.clone(), password).await? {
280        Some(token) => token,
281        None => return Ok(None),
282    };
283    debug!("Successfully obtained OAuth token from TMC");
284
285    let tmc_user = tmc_client
286        .get_user_from_tmc_mooc_fi_by_tmc_access_token(&token.clone())
287        .await?;
288    debug!(
289        "Creating or fetching user with TMC id {} and mooc.fi UUID {}",
290        tmc_user.id,
291        tmc_user
292            .courses_mooc_fi_user_id
293            .map(|uuid| uuid.to_string())
294            .unwrap_or_else(|| "None (will fetch from mooc.fi or generate new UUID)".to_string())
295    );
296    let user = get_or_create_user_from_tmc_mooc_fi_response(&mut *conn, tmc_user, &token).await?;
297    info!(
298        "Successfully got user details from mooc.fi for user {}",
299        user.id
300    );
301    info!("Successfully authenticated user {} with mooc.fi", user.id);
302    Ok(Some((user, token)))
303}
304
305pub type LoginToken = StandardTokenResponse<EmptyExtraTokenFields, BasicTokenType>;
306
307/// Exchanges user credentials with TMC for an OAuth token.
308///
309/// `Ok(None)` means the credentials were rejected; other failures (network, server errors) are
310/// `Err`.
311pub async fn exchange_password_with_tmc(
312    client: &OAuthClient,
313    email: String,
314    password: SecretString,
315) -> anyhow::Result<Option<SecretString>> {
316    let token_result = client
317        .exchange_password(
318            &ResourceOwnerUsername::new(email),
319            // Exposed only here, at the OAuth2 client boundary.
320            &ResourceOwnerPassword::new(password.expose_secret().to_string()),
321        )
322        .request_async(&async_http_client_with_headers)
323        .await;
324    match token_result {
325        Ok(token) => Ok(Some(SecretString::new(
326            token.access_token().secret().to_owned().into(),
327        ))),
328        Err(RequestTokenError::ServerResponse(server_response)) => {
329            let error = server_response.error();
330            let error_description = server_response.error_description();
331            let error_uri = server_response.error_uri();
332
333            // InvalidGrant means the email or password was wrong.
334            if let oauth2::basic::BasicErrorResponseType::InvalidGrant = error {
335                warn!(
336                    ?error_description,
337                    ?error_uri,
338                    "TMC did not accept the credentials: {}",
339                    error
340                );
341                Ok(None)
342            } else {
343                error!(
344                    ?error_description,
345                    ?error_uri,
346                    "TMC authentication error: {}",
347                    error
348                );
349                Err(anyhow::anyhow!("Authentication error: {}", error))
350            }
351        }
352        Err(e) => {
353            error!("Failed to exchange password with TMC: {}", e);
354            Err(e.into())
355        }
356    }
357}
358
359/// Fetches the mooc.fi UUID for a user by their upstream ID using the TMC access token.
360async fn fetch_moocfi_id_by_upstream_id(
361    tmc_access_token: &SecretString,
362    upstream_id: i32,
363) -> anyhow::Result<Option<Uuid>> {
364    info!("Fetching mooc.fi UUID for upstream user id {}", upstream_id);
365
366    let res = REQWEST_CLIENT
367        .post(MOOCFI_GRAPHQL_URL)
368        .header(reqwest::header::CONTENT_TYPE, "application/json")
369        .header(reqwest::header::ACCEPT, "application/json")
370        // Exposed only here, where the bearer token header is built.
371        .bearer_auth(tmc_access_token.expose_secret())
372        .json(&GraphQLRequest {
373            query: r#"
374query ($upstreamId: Int) {
375  user(upstream_id: $upstreamId) {
376    id
377  }
378}"#,
379            variables: Some(json!({ "upstreamId": upstream_id })),
380        })
381        .send()
382        .await;
383
384    match res {
385        Ok(response) => {
386            if !response.status().is_success() {
387                debug!(
388                    "Failed to fetch mooc.fi user with status {}. Will generate new UUID instead.",
389                    response.status()
390                );
391                return Ok(None);
392            }
393
394            match response.json::<MoocfiUserResponse>().await {
395                Ok(current_user_response) => {
396                    info!(
397                        "Successfully fetched mooc.fi UUID {} for upstream id {}",
398                        current_user_response.data.user.id, upstream_id
399                    );
400                    Ok(Some(current_user_response.data.user.id))
401                }
402                Err(e) => {
403                    debug!(
404                        "Failed to parse mooc.fi response: {}. Will generate new UUID instead.",
405                        e
406                    );
407                    Ok(None)
408                }
409            }
410        }
411        Err(e) => {
412            debug!(
413                "Failed to fetch from mooc.fi: {}. Will generate new UUID instead.",
414                e
415            );
416            Ok(None)
417        }
418    }
419}
420
421pub async fn get_or_create_user_from_tmc_mooc_fi_response(
422    conn: &mut PgConnection,
423    tmc_mooc_fi_user: TMCUser,
424    tmc_access_token: &SecretString,
425) -> anyhow::Result<User> {
426    let TMCUser {
427        id: upstream_id,
428        email,
429        courses_mooc_fi_user_id: moocfi_id,
430        user_field,
431        ..
432    } = tmc_mooc_fi_user;
433
434    let id = match moocfi_id {
435        Some(id) => id,
436        None => match fetch_moocfi_id_by_upstream_id(tmc_access_token, upstream_id).await? {
437            Some(fetched_id) => {
438                info!("Successfully fetched mooc.fi UUID {} for user", fetched_id);
439                fetched_id
440            }
441            None => {
442                info!("No mooc.fi UUID found, generating new UUID for user");
443                Uuid::new_v4()
444            }
445        },
446    };
447
448    let user = match models::users::find_by_upstream_id(conn, upstream_id).await? {
449        Some(existing_user) => existing_user,
450        None => {
451            let inserted = models::users::insert_with_upstream_id_and_moocfi_id(
452                conn,
453                &email,
454                user_field
455                    .first_name
456                    .as_deref()
457                    .filter(|s| !s.trim().is_empty()),
458                user_field
459                    .last_name
460                    .as_deref()
461                    .filter(|s| !s.trim().is_empty()),
462                upstream_id,
463                id,
464            )
465            .await;
466            match inserted {
467                Ok(user) => user,
468                // A concurrent request can create the user between the find and the insert
469                // (the insert runs in a savepoint, so the connection stays usable). The unique
470                // index on upstream_id rejects the loser; return the winner's row instead.
471                Err(insert_error)
472                    if matches!(
473                        insert_error.error_type(),
474                        models::ModelErrorType::DatabaseConstraint { constraint, .. }
475                            if constraint == "users_upstream_id_active_uniq_idx"
476                    ) =>
477                {
478                    models::users::find_by_upstream_id(conn, upstream_id)
479                        .await?
480                        .ok_or(insert_error)?
481                }
482                Err(insert_error) => return Err(insert_error.into()),
483            }
484        }
485    };
486    Ok(user)
487}
488
489/// Authenticates a test user against predefined credentials.
490pub async fn authenticate_test_user(
491    conn: &mut PgConnection,
492    email: &str,
493    password: &SecretString,
494    application_configuration: &ApplicationConfiguration,
495) -> anyhow::Result<bool> {
496    // Sanity check to ensure this is not called outside of test mode. The whole application configuration is passed to this function instead of just the boolean to make mistakes harder.
497    assert!(application_configuration.test_mode);
498
499    // Test-only seeded credentials; exposed once here for the literal comparisons below.
500    let password = password.expose_secret();
501
502    let _user = if email == "admin@example.com" && password == "admin" {
503        models::users::get_by_email(conn, "admin@example.com").await?
504    } else if email == "teacher@example.com" && password == "teacher" {
505        models::users::get_by_email(conn, "teacher@example.com").await?
506    } else if email == "language.teacher@example.com" && password == "language.teacher" {
507        models::users::get_by_email(conn, "language.teacher@example.com").await?
508    } else if email == "material.viewer@example.com" && password == "material.viewer" {
509        models::users::get_by_email(conn, "material.viewer@example.com").await?
510    } else if email == "user@example.com" && password == "user" {
511        models::users::get_by_email(conn, "user@example.com").await?
512    } else if email == "assistant@example.com" && password == "assistant" {
513        models::users::get_by_email(conn, "assistant@example.com").await?
514    } else if email == "creator@example.com" && password == "creator" {
515        models::users::get_by_email(conn, "creator@example.com").await?
516    } else if email == "student1@example.com" && password == "student1" {
517        models::users::get_by_email(conn, "student1@example.com").await?
518    } else if email == "student2@example.com" && password == "student2" {
519        models::users::get_by_email(conn, "student2@example.com").await?
520    } else if email == "student3@example.com" && password == "student3" {
521        models::users::get_by_email(conn, "student3@example.com").await?
522    } else if email == "student4@example.com" && password == "student4" {
523        models::users::get_by_email(conn, "student4@example.com").await?
524    } else if email == "student5@example.com" && password == "student5" {
525        models::users::get_by_email(conn, "student5@example.com").await?
526    } else if email == "student6@example.com" && password == "student6" {
527        models::users::get_by_email(conn, "student6@example.com").await?
528    } else if email == "student7@example.com" && password == "student7" {
529        models::users::get_by_email(conn, "student7@example.com").await?
530    } else if email == "student8@example.com" && password == "student8" {
531        models::users::get_by_email(conn, "student8@example.com").await?
532    } else if email == "teaching-and-learning-services@example.com"
533        && password == "teaching-and-learning-services"
534    {
535        models::users::get_by_email(conn, "teaching-and-learning-services@example.com").await?
536    } else if email == "student-without-research-consent@example.com"
537        && password == "student-without-research-consent"
538    {
539        models::users::get_by_email(conn, "student-without-research-consent@example.com").await?
540    } else if email == "student-without-country@example.com"
541        && password == "student-without-country"
542    {
543        models::users::get_by_email(conn, "student-without-country@example.com").await?
544    } else if email == "langs@example.com" && password == "langs" {
545        models::users::get_by_email(conn, "langs@example.com").await?
546    } else if email == "sign-up-user@example.com" && password == "sign-up-user" {
547        models::users::get_by_email(conn, "sign-up-user@example.com").await?
548    } else {
549        info!("Authentication failed: incorrect test credentials");
550        return Ok(false);
551    };
552    info!("Successfully authenticated test user {}", email);
553    Ok(true)
554}
555
556// Only used for testing, not to use in production.
557pub async fn authenticate_test_token(
558    conn: &mut PgConnection,
559    token: &SecretString,
560    application_configuration: &ApplicationConfiguration,
561) -> anyhow::Result<Option<User>> {
562    // Sanity check to ensure this is not called outside of test mode. The whole application configuration is passed to this function instead of just the boolean to make mistakes harder.
563    assert!(application_configuration.test_mode);
564
565    // These token strings are well-known constants, not secrets; they only work under
566    // `test_mode`.
567    let email = match token.expose_secret() {
568        "test-token-langs" => "langs@example.com",
569        "test-token-student1" => "student1@example.com",
570        "test-token-student2" => "student2@example.com",
571        _ => return Ok(None),
572    };
573    let user = models::users::get_by_email(conn, email).await?;
574    info!("Test mode: mapped fixed test token to seeded user {email}");
575    Ok(Some(user))
576}
577
578/// The rate-limit-bypass header value for requests to the TMC server.
579fn get_ratelimit_api_key() -> Result<reqwest::header::HeaderValue, HttpClientError<reqwest::Error>>
580{
581    let key = server_runtime_config()
582        .ratelimit_protection_safe_api_key
583        .clone();
584    debug!("Using ratelimit API key from runtime config");
585
586    key.expose_secret()
587        .parse::<reqwest::header::HeaderValue>()
588        .map_err(|err| {
589            error!("Invalid RATELIMIT API key format: {}", err);
590            HttpClientError::Other("Invalid RATELIMIT API key.".to_string())
591        })
592}
593
594/// oauth2's HTTP transport, adapted to reqwest and tagged with the header that keeps TMC from
595/// rate-limiting the backend's own auth requests.
596async fn async_http_client_with_headers(
597    oauth_request: oauth2::HttpRequest,
598) -> Result<oauth2::HttpResponse, HttpClientError<reqwest::Error>> {
599    debug!("Making OAuth request to TMC server");
600
601    if log::log_enabled!(log::Level::Trace) {
602        // Only log the URL path, not query parameters which may contain credentials
603        if let Ok(url) = oauth_request.uri().to_string().parse::<reqwest::Url>() {
604            trace!("OAuth request path: {}", url.path());
605        }
606    }
607
608    let parsed_key = get_ratelimit_api_key()?;
609
610    debug!("Building request to TMC server");
611    let request = REQWEST_CLIENT
612        .request(
613            oauth_request.method().clone(),
614            oauth_request
615                .uri()
616                .to_string()
617                .parse::<reqwest::Url>()
618                .map_err(|e| HttpClientError::Other(format!("Invalid URL: {}", e)))?,
619        )
620        .headers(oauth_request.headers().clone())
621        .version(oauth_request.version())
622        .header("RATELIMIT-PROTECTION-SAFE-API-KEY", parsed_key)
623        .body(oauth_request.body().to_vec());
624
625    debug!("Sending request to TMC server");
626    let response = request
627        .send()
628        .await
629        .map_err(|e| HttpClientError::Other(format!("Failed to execute request: {}", e)))?;
630
631    // Log response status and version, but not headers or body which may contain tokens
632    debug!(
633        "Received response from TMC server - Status: {}, Version: {:?}",
634        response.status(),
635        response.version()
636    );
637
638    let status = response.status();
639    let version = response.version();
640    let headers = response.headers().clone();
641
642    debug!("Reading response body");
643    let body_bytes = response
644        .bytes()
645        .await
646        .map_err(|e| HttpClientError::Other(format!("Failed to read response body: {}", e)))?
647        .to_vec();
648
649    debug!("Building OAuth response");
650    let mut builder = oauth2::http::Response::builder()
651        .status(status)
652        .version(version);
653
654    if let Some(builder_headers) = builder.headers_mut() {
655        builder_headers.extend(headers.iter().map(|(k, v)| (k.clone(), v.clone())));
656    }
657
658    let oauth_response = builder
659        .body(body_bytes)
660        .map_err(|e| HttpClientError::Other(format!("Failed to construct response: {}", e)))?;
661
662    debug!("Successfully completed OAuth request");
663    Ok(oauth_response)
664}