1use 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#[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 pub fn upstream_id(&self) -> Option<i32> {
80 self.upstream_id
81 }
82}
83
84pub 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.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
123async 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
166pub 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
181pub 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
193pub fn forget(session: &Session) {
195 session.purge();
196}
197
198pub 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
213pub 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
270pub 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
307pub 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 &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 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
359async 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 .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 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
489pub async fn authenticate_test_user(
491 conn: &mut PgConnection,
492 email: &str,
493 password: &SecretString,
494 application_configuration: &ApplicationConfiguration,
495) -> anyhow::Result<bool> {
496 assert!(application_configuration.test_mode);
498
499 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
556pub async fn authenticate_test_token(
558 conn: &mut PgConnection,
559 token: &SecretString,
560 application_configuration: &ApplicationConfiguration,
561) -> anyhow::Result<Option<User>> {
562 assert!(application_configuration.test_mode);
564
565 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
578fn 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
594async 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 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 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}