Skip to main content

headless_lms_server/domain/
models_requests.rs

1//! Contains helper functions that are passed to headless-lms-models where it needs to make requests to exercise services.
2
3use crate::prelude::*;
4use actix_http::Payload;
5use actix_web::{FromRequest, HttpRequest};
6use chrono::{Duration, Utc};
7use futures::{
8    FutureExt,
9    future::{BoxFuture, Ready, ready},
10};
11use headless_lms_models::{
12    HttpErrorType, ModelError, ModelErrorType, ModelResult,
13    exercise_service_info::ExerciseServiceInfoApi,
14    exercise_task_gradings::{
15        ExerciseTaskGradingRequest, ExerciseTaskGradingResult, GradingRequestFile,
16    },
17    exercise_task_submissions::{AnswerFile as SubmittedAnswerFile, ExerciseTaskSubmission},
18    exercise_tasks::ExerciseTask,
19};
20
21use headless_lms_base::error::backend_error::BackendError;
22pub use headless_lms_base::jwt::{DOWNLOAD_CLAIM_PARAM, DownloadClaim, JwtKey};
23use headless_lms_base::jwt::{claimed_file_url, sign_hs256_claim, validate_hs256_claim};
24use models::SpecFetcher;
25use std::collections::HashMap;
26use std::fmt::Debug;
27use std::sync::{Arc, Mutex};
28use url::Url;
29
30use super::error::{ControllerError, ControllerErrorType};
31
32// keep in sync with the shared-module constants
33const EXERCISE_SERVICE_GRADING_UPDATE_CLAIM_HEADER: &str = "exercise-service-grading-update-claim";
34const EXERCISE_SERVICE_UPLOAD_CLAIM_HEADER: &str = "exercise-service-upload-claim";
35pub const PLAYGROUND_GRADING_CALLBACK_CLAIM_PARAM: &str = "playground-grading-callback-claim";
36
37/// A type for caching the spec fetching (only for the seed)
38type SpecCache = HashMap<(String, String, Option<String>), serde_json::Value>;
39
40#[derive(Debug, Serialize, Deserialize)]
41pub struct UploadClaim {
42    exercise_service_slug: String,
43    exp: usize,
44    iat: usize,
45}
46
47impl UploadClaim {
48    pub fn exercise_service_slug(&self) -> &str {
49        self.exercise_service_slug.as_ref()
50    }
51
52    pub fn expiring_in_1_day(exercise_service_slug: impl Into<String>) -> Self {
53        let now = Utc::now().timestamp().max(0) as usize;
54        let exp = (Utc::now().timestamp() + Duration::days(1).num_seconds()).max(0) as usize;
55        Self {
56            exercise_service_slug: exercise_service_slug.into(),
57            exp,
58            iat: now,
59        }
60    }
61
62    pub fn sign(self, key: &JwtKey) -> Result<String, jsonwebtoken::errors::Error> {
63        sign_hs256_claim(&self, key)
64    }
65
66    pub fn validate(token: &str, key: &JwtKey) -> Result<Self, ControllerError> {
67        validate_claim(token, key)
68    }
69}
70
71impl FromRequest for UploadClaim {
72    type Error = ControllerError;
73    type Future = Ready<Result<Self, Self::Error>>;
74
75    fn from_request(req: &HttpRequest, _payload: &mut Payload) -> Self::Future {
76        let try_from_request = move || {
77            let jwt_key = req.app_data::<web::Data<JwtKey>>().ok_or_else(|| {
78                ControllerError::new(
79                    ControllerErrorType::InternalServerError,
80                    "Missing JwtKey in app data - server configuration error".to_string(),
81                    None,
82                )
83            })?;
84            let header = req
85                .headers()
86                .get(EXERCISE_SERVICE_UPLOAD_CLAIM_HEADER)
87                .ok_or_else(|| {
88                    ControllerError::new(
89                        ControllerErrorType::BadRequest,
90                        format!("Missing header {EXERCISE_SERVICE_UPLOAD_CLAIM_HEADER}",),
91                        None,
92                    )
93                })?;
94            let header = std::str::from_utf8(header.as_bytes()).map_err(|err| {
95                ControllerError::new(
96                    ControllerErrorType::BadRequest,
97                    format!(
98                        "Invalid header {EXERCISE_SERVICE_UPLOAD_CLAIM_HEADER} = {}",
99                        String::from_utf8_lossy(header.as_bytes())
100                    ),
101                    Some(err.into()),
102                )
103            })?;
104            let claim = UploadClaim::validate(header, jwt_key)?;
105            Result::<_, Self::Error>::Ok(claim)
106        };
107        ready(try_from_request())
108    }
109}
110
111#[derive(Debug, Serialize, Deserialize)]
112pub struct GradingUpdateClaim {
113    submission_id: Uuid,
114    exp: usize,
115    iat: usize,
116}
117
118impl GradingUpdateClaim {
119    pub fn submission_id(&self) -> Uuid {
120        self.submission_id
121    }
122
123    pub fn expiring_in_1_day(submission_id: Uuid) -> Self {
124        let now = Utc::now().timestamp().max(0) as usize;
125        let exp = (Utc::now().timestamp() + Duration::days(1).num_seconds()).max(0) as usize;
126        Self {
127            submission_id,
128            exp,
129            iat: now,
130        }
131    }
132
133    pub fn sign(self, key: &JwtKey) -> Result<String, jsonwebtoken::errors::Error> {
134        sign_hs256_claim(&self, key)
135    }
136
137    pub fn validate(token: &str, key: &JwtKey) -> Result<Self, ControllerError> {
138        validate_claim(token, key)
139    }
140}
141
142impl FromRequest for GradingUpdateClaim {
143    type Error = ControllerError;
144    type Future = Ready<Result<Self, Self::Error>>;
145
146    fn from_request(req: &HttpRequest, _payload: &mut Payload) -> Self::Future {
147        let try_from_request = move || {
148            let jwt_key = req.app_data::<web::Data<JwtKey>>().ok_or_else(|| {
149                ControllerError::new(
150                    ControllerErrorType::InternalServerError,
151                    "Missing JwtKey in app data - server configuration error".to_string(),
152                    None,
153                )
154            })?;
155            let header = req
156                .headers()
157                .get(EXERCISE_SERVICE_GRADING_UPDATE_CLAIM_HEADER)
158                .ok_or_else(|| {
159                    ControllerError::new(
160                        ControllerErrorType::BadRequest,
161                        format!("Missing header {EXERCISE_SERVICE_GRADING_UPDATE_CLAIM_HEADER}",),
162                        None,
163                    )
164                })?;
165            let header = std::str::from_utf8(header.as_bytes()).map_err(|err| {
166                ControllerError::new(
167                    ControllerErrorType::BadRequest,
168                    format!(
169                        "Invalid header {EXERCISE_SERVICE_GRADING_UPDATE_CLAIM_HEADER} = {}",
170                        String::from_utf8_lossy(header.as_bytes())
171                    ),
172                    Some(err.into()),
173                )
174            })?;
175            let claim = GradingUpdateClaim::validate(header, jwt_key)?;
176            Result::<_, Self::Error>::Ok(claim)
177        };
178        ready(try_from_request())
179    }
180}
181
182#[derive(Debug, Serialize, Deserialize)]
183pub struct PlaygroundGradingCallbackClaim {
184    websocket_id: Uuid,
185    exp: usize,
186    iat: usize,
187}
188
189impl PlaygroundGradingCallbackClaim {
190    pub fn websocket_id(&self) -> Uuid {
191        self.websocket_id
192    }
193
194    pub fn expiring_in_1_day(websocket_id: Uuid) -> Self {
195        let now = Utc::now().timestamp().max(0) as usize;
196        let exp = (Utc::now().timestamp() + Duration::days(1).num_seconds()).max(0) as usize;
197        Self {
198            websocket_id,
199            exp,
200            iat: now,
201        }
202    }
203
204    pub fn sign(self, key: &JwtKey) -> Result<String, jsonwebtoken::errors::Error> {
205        sign_hs256_claim(&self, key)
206    }
207
208    pub fn validate(token: &str, key: &JwtKey) -> Result<Self, ControllerError> {
209        validate_hs256_claim::<Self>(token, key).map_err(|err| {
210            controller_err!(
211                BadRequest,
212                format!("Invalid playground grading callback claim: {}", err),
213                err
214            )
215        })
216    }
217}
218
219impl FromRequest for PlaygroundGradingCallbackClaim {
220    type Error = ControllerError;
221    type Future = Ready<Result<Self, Self::Error>>;
222
223    fn from_request(req: &HttpRequest, _payload: &mut Payload) -> Self::Future {
224        let try_from_request = move || {
225            let jwt_key = req.app_data::<web::Data<JwtKey>>().ok_or_else(|| {
226                controller_err!(
227                    InternalServerError,
228                    "Missing JwtKey in app data - server configuration error".to_string()
229                )
230            })?;
231            let query_claim = url::form_urlencoded::parse(req.query_string().as_bytes())
232                .find(|(key, _)| key == PLAYGROUND_GRADING_CALLBACK_CLAIM_PARAM)
233                .map(|(_, value)| value.into_owned());
234            let header_claim = req
235                .headers()
236                .get(PLAYGROUND_GRADING_CALLBACK_CLAIM_PARAM)
237                .and_then(|header| std::str::from_utf8(header.as_bytes()).ok())
238                .map(ToString::to_string);
239            let claim = header_claim.or(query_claim).ok_or_else(|| {
240                controller_err!(
241                    BadRequest,
242                    format!("Missing {PLAYGROUND_GRADING_CALLBACK_CLAIM_PARAM}")
243                )
244            })?;
245            PlaygroundGradingCallbackClaim::validate(&claim, jwt_key)
246        };
247        ready(try_from_request())
248    }
249}
250
251/// Accepted by the public-spec and model-solution endpoints of exercise services.
252#[derive(Debug, Serialize)]
253
254pub struct SpecRequest<'a> {
255    request_id: Uuid,
256    private_spec: Option<&'a serde_json::Value>,
257    upload_url: Option<String>,
258}
259
260#[derive(Debug, Serialize)]
261pub struct ExerciseServiceCsvExportRequest<'a, T: Serialize> {
262    pub items: &'a [T],
263}
264
265/// Column definition for exercise service CSV export; callers must use scalar-only cell values.
266#[derive(Debug, Serialize, Deserialize, Clone, PartialEq)]
267pub struct ExerciseServiceCsvExportColumn {
268    pub key: String,
269    pub header: String,
270}
271
272/// One batch of CSV rows; each row's values must be scalar (null, bool, number, string). Objects/arrays are rejected by the controller.
273#[derive(Debug, Serialize, Deserialize, Clone, PartialEq)]
274pub struct ExerciseServiceCsvExportResult {
275    pub rows: Vec<HashMap<String, serde_json::Value>>,
276}
277
278/// Full CSV export response; columns define headers, results align by index. All cell values must be scalar.
279#[derive(Debug, Serialize, Deserialize, Clone, PartialEq)]
280pub struct ExerciseServiceCsvExportResponse {
281    pub columns: Vec<ExerciseServiceCsvExportColumn>,
282    pub results: Vec<ExerciseServiceCsvExportResult>,
283}
284
285/// Fetches a public/model spec based on the private spec from the given url.
286/// The slug and jwt key are used for an upload claim that allows the service
287/// to upload files as part of the spec.
288pub fn make_spec_fetcher(
289    base_url: String,
290    request_id: Uuid,
291    jwt_key: Arc<JwtKey>,
292) -> impl SpecFetcher {
293    move |url, exercise_service_slug, private_spec| {
294        let client = reqwest::Client::new();
295        let upload_claim = UploadClaim::expiring_in_1_day(exercise_service_slug);
296        let upload_url = Some(format!("{base_url}/api/v0/files/{exercise_service_slug}"));
297        let signed_upload_claim = match upload_claim.sign(&jwt_key) {
298            Ok(claim) => claim,
299            Err(err) => {
300                return async move {
301                    Err(ModelError::new(
302                        ModelErrorType::Generic,
303                        format!("Failed to sign upload claim: {err}"),
304                        Some(err.into()),
305                    ))
306                }
307                .boxed();
308            }
309        };
310        let req = client
311            .post(url.clone())
312            .header(EXERCISE_SERVICE_UPLOAD_CLAIM_HEADER, signed_upload_claim)
313            .timeout(std::time::Duration::from_secs(120))
314            .json(&SpecRequest {
315                request_id,
316                private_spec,
317                upload_url,
318            })
319            .send();
320        async move {
321            let res = req.await.map_err(ModelError::from)?;
322            let status_code = res.status();
323            if !status_code.is_success() {
324                let error_text = res.text().await;
325                let error = error_text.as_deref().unwrap_or("(No text in response)");
326                error!(
327                    ?url,
328                    ?exercise_service_slug,
329                    ?private_spec,
330                    ?status_code,
331                    "Exercise service returned an error while generating a spec: {}",
332                    error
333                );
334                return Err(ModelError::new(
335                    ModelErrorType::HttpRequest {
336                        status_code: status_code.as_u16(),
337                        response_body: error.to_string(),
338                    },
339                    format!(
340                        "Failed to generate spec for exercise for {exercise_service_slug}: {error}."
341                    ),
342                    None,
343                ));
344            }
345            let json = parse_response_json(res).await?;
346            Ok(json)
347        }
348        .boxed()
349    }
350}
351
352// see `fetch_service_info_fast` while handling HTTP requests
353pub fn fetch_service_info(url: Url) -> BoxFuture<'static, ModelResult<ExerciseServiceInfoApi>> {
354    fetch_service_info_with_timeout(url, 1000 * 120)
355}
356
357// use this while handling HTTP requests, see `fetch_service_info`
358pub fn fetch_service_info_fast(
359    url: Url,
360) -> BoxFuture<'static, ModelResult<ExerciseServiceInfoApi>> {
361    fetch_service_info_with_timeout(url, 1000 * 5)
362}
363
364fn fetch_service_info_with_timeout(
365    url: Url,
366    timeout_ms: u64,
367) -> BoxFuture<'static, ModelResult<ExerciseServiceInfoApi>> {
368    async move {
369        let client = reqwest::Client::new();
370        let res = client
371            .get(url) // e.g. http://example-exercise.default.svc.cluster.local:3002/example-exercise/api/service-info
372            .timeout(std::time::Duration::from_millis(timeout_ms))
373            .send()
374            .await
375            .map_err(ModelError::from)?;
376        let status = res.status();
377        if !status.is_success() {
378            let response_url = res.url().to_string();
379            let body = res.text().await.map_err(ModelError::from)?;
380            warn!(url=?response_url, status=?status, body=?body, "Could not fetch service info.");
381            return Err(ModelError::new(
382                ModelErrorType::HttpRequest {
383                    status_code: status.as_u16(),
384                    response_body: body,
385                },
386                "Could not fetch service info.".to_string(),
387                None,
388            ));
389        }
390        let res = parse_response_json(res).await?;
391        Ok(res)
392    }
393    .boxed()
394}
395
396/// The grading request's file list for a submission's answer: its files in answer order, empty for
397/// an answer that has none.
398///
399/// Mints a single-file download claim per file, so the service fetches a URL the host chose rather
400/// than one a student supplied.
401fn grading_request_files(
402    files: Option<&[SubmittedAnswerFile]>,
403    base_url: &str,
404    jwt_key: &JwtKey,
405) -> Result<Vec<GradingRequestFile>, jsonwebtoken::errors::Error> {
406    let Some(files) = files else {
407        return Ok(Vec::new());
408    };
409    let mut ordered: Vec<&SubmittedAnswerFile> = files.iter().collect();
410    ordered.sort_by_key(|file| file.order_number);
411    ordered
412        .into_iter()
413        .map(|file| {
414            Ok(GradingRequestFile {
415                id: file.id,
416                name: file.name.clone(),
417                mime: file.mime.clone(),
418                size_bytes: file.size_bytes,
419                download_url: claimed_file_url(
420                    base_url,
421                    jwt_key,
422                    DownloadClaim::expiring_in_1_day(file.id),
423                )?,
424            })
425        })
426        .collect()
427}
428
429/// Sends a submission to an exercise service for grading, with the claims the service needs to
430/// report back and to read a file-typed answer's files.
431///
432/// `base_url` is the host's own public base url, so both the callback and the file urls point at a
433/// host the service can reach.
434pub fn make_grading_request_sender(
435    jwt_key: Arc<JwtKey>,
436    base_url: String,
437) -> impl Fn(
438    Url,
439    &ExerciseTask,
440    &ExerciseTaskSubmission,
441) -> BoxFuture<'static, ModelResult<ExerciseTaskGradingResult>> {
442    move |grade_url, exercise_task, submission| {
443        let client = reqwest::Client::new();
444        let grading_update_url = format!(
445            "{base_url}/api/v0/exercise-services/grading/grading-update/{}",
446            submission.id
447        );
448        let submission_files =
449            match grading_request_files(submission.data_files.as_deref(), &base_url, &jwt_key) {
450                Ok(files) => files,
451                Err(err) => {
452                    return async move {
453                        Err(ModelError::new(
454                            ModelErrorType::Generic,
455                            format!("Failed to sign download claim: {err}"),
456                            Some(err.into()),
457                        ))
458                    }
459                    .boxed();
460                }
461            };
462        let grading_update_claim = GradingUpdateClaim::expiring_in_1_day(submission.id);
463        let signed_grading_update_claim = match grading_update_claim.sign(&jwt_key) {
464            Ok(claim) => claim,
465            Err(err) => {
466                return async move {
467                    Err(ModelError::new(
468                        ModelErrorType::Generic,
469                        format!("Failed to sign grading update claim: {err}"),
470                        Some(err.into()),
471                    ))
472                }
473                .boxed();
474            }
475        };
476        let req = client
477            .post(grade_url)
478            .header(
479                EXERCISE_SERVICE_GRADING_UPDATE_CLAIM_HEADER,
480                signed_grading_update_claim,
481            )
482            .timeout(std::time::Duration::from_secs(120))
483            .json(&ExerciseTaskGradingRequest {
484                grading_update_url: &grading_update_url,
485                exercise_spec: &exercise_task.private_spec,
486                submission_data: submission.data_json.as_ref(),
487                submission_files: &submission_files,
488            });
489        async move {
490            let res = req.send().await.map_err(ModelError::from)?;
491            let status = res.status();
492            if !status.is_success() {
493                let status_code = status.as_u16();
494                let response_body = res.text().await.unwrap_or_default();
495                error!(
496                    ?response_body,
497                    status_code = %status_code,
498                    "Grading request returned an unsuccesful status code"
499                );
500
501                return Err(ModelError::new(
502                    ModelErrorType::HttpRequest {
503                        status_code,
504                        response_body: response_body.clone(),
505                    },
506                    format!(
507                        "Grading failed with status: {} response: {}",
508                        status_code, response_body
509                    ),
510                    None,
511                ));
512            }
513            let obj = parse_response_json(res).await?;
514            info!("Received a grading result: {:#?}", &obj);
515            Ok(obj)
516        }
517        .boxed()
518    }
519}
520
521pub async fn post_exercise_service_csv_export_request<T: Serialize>(
522    url: Url,
523    items: &[T],
524) -> ModelResult<ExerciseServiceCsvExportResponse> {
525    let client = reqwest::Client::new();
526    let response = client
527        .post(url.clone())
528        .timeout(std::time::Duration::from_secs(120))
529        .json(&ExerciseServiceCsvExportRequest { items })
530        .send()
531        .await
532        .map_err(ModelError::from)?;
533
534    let status = response.status();
535    if !status.is_success() {
536        let status_code = status.as_u16();
537        let response_body = response.text().await.unwrap_or_default();
538        error!(
539            ?response_body,
540            status_code = %status_code,
541            "Exercise service CSV export request returned an unsuccessful status code"
542        );
543
544        return Err(ModelError::new(
545            ModelErrorType::HttpRequest {
546                status_code,
547                response_body: response_body.clone(),
548            },
549            format!(
550                "CSV export request failed with status: {} response: {}",
551                status_code, response_body
552            ),
553            None,
554        ));
555    }
556
557    parse_response_json(response).await
558}
559
560#[derive(Debug, Serialize, Deserialize)]
561pub struct GivePeerReviewClaim {
562    pub exercise_slide_submission_id: Uuid,
563    pub peer_or_self_review_config_id: Uuid,
564    exp: usize,
565    iat: usize,
566}
567
568impl GivePeerReviewClaim {
569    pub fn expiring_in_1_day(
570        exercise_slide_submission_id: Uuid,
571        peer_or_self_review_config_id: Uuid,
572    ) -> Self {
573        let now = Utc::now().timestamp().max(0) as usize;
574        let exp = (Utc::now().timestamp() + Duration::days(1).num_seconds()).max(0) as usize;
575        Self {
576            exercise_slide_submission_id,
577            peer_or_self_review_config_id,
578            exp,
579            iat: now,
580        }
581    }
582
583    pub fn sign(self, key: &JwtKey) -> Result<String, jsonwebtoken::errors::Error> {
584        sign_hs256_claim(&self, key)
585    }
586
587    pub fn validate(token: &str, key: &JwtKey) -> Result<Self, ControllerError> {
588        validate_hs256_claim(token, key).map_err(|err| {
589            ControllerError::new(
590                ControllerErrorType::BadRequest,
591                format!("Invalid claim: {}", err),
592                Some(err.into()),
593            )
594        })
595    }
596}
597
598/// Decodes a claim, reporting a bad token as a request error rather than a JWT one.
599///
600/// [`validate_hs256_claim`] is the raw form that leaves the `jsonwebtoken` error unmapped, for the
601/// claims that report it differently.
602fn validate_claim<T: serde::de::DeserializeOwned>(
603    token: &str,
604    key: &JwtKey,
605) -> Result<T, ControllerError> {
606    validate_hs256_claim(token, key)
607        .map_err(|err| controller_err!(BadRequest, format!("Invalid jwt key: {}", err), err))
608}
609
610/// A caching spec fetcher ONLY FOR THE SEED that returns a cached spec if the same
611/// (url, exercise_service_slug, private_spec) is requested. Since this is only used during seeding,
612/// there is no cache eviction.
613pub fn make_seed_spec_fetcher_with_cache(
614    base_url: String,
615    request_id: Uuid,
616    jwt_key: Arc<JwtKey>,
617) -> impl SpecFetcher {
618    // Cache key: (url, exercise_service_slug, private_spec serialized)
619    let cache: Arc<Mutex<SpecCache>> = Arc::new(Mutex::new(HashMap::new()));
620
621    // Create the base non-caching spec fetcher and wrap it in Arc to make it clonable
622    let base_fetcher = Arc::new(make_spec_fetcher(base_url, request_id, jwt_key));
623
624    move |url, exercise_service_slug, private_spec| {
625        let url_str = url.to_string();
626        let service_slug = exercise_service_slug.to_string();
627        // Convert private_spec to string for cache key if present
628        let private_spec_str =
629            private_spec.map(|spec| serde_json::to_string(&spec).unwrap_or_default());
630        let key = (url_str.clone(), service_slug.clone(), private_spec_str);
631        let cache = Arc::clone(&cache);
632        let base_fetcher = Arc::clone(&base_fetcher);
633
634        async move {
635            // Try to get from cache first
636            let cached_spec = {
637                let cache_guard = cache.lock().map_err(|err| {
638                    ModelError::new(
639                        ModelErrorType::Generic,
640                        format!("Seed spec fetcher cache lock poisoned: {err}"),
641                        None::<anyhow::Error>,
642                    )
643                })?;
644                cache_guard.get(&key).cloned()
645            };
646            if let Some(cached_spec) = cached_spec {
647                return Ok(cached_spec.clone());
648            }
649
650            // Not in cache - fetch using base fetcher
651            let fetched_spec = base_fetcher(url, exercise_service_slug, private_spec).await?;
652
653            // Store in cache
654            {
655                let mut cache_guard = cache.lock().map_err(|err| {
656                    ModelError::new(
657                        ModelErrorType::Generic,
658                        format!("Seed spec fetcher cache lock poisoned: {err}"),
659                        None::<anyhow::Error>,
660                    )
661                })?;
662                cache_guard.insert(key, fetched_spec.clone());
663            }
664
665            Ok(fetched_spec)
666        }
667        .boxed()
668    }
669}
670
671/// Safely parses a response body as JSON, capturing the actual response body in error cases
672async fn parse_response_json<T>(response: reqwest::Response) -> ModelResult<T>
673where
674    T: serde::de::DeserializeOwned,
675{
676    let status = response.status();
677    let response_text = response.text().await.map_err(ModelError::from)?;
678
679    serde_json::from_str(&response_text).map_err(|err| {
680        ModelError::new(
681            ModelErrorType::HttpError {
682                error_type: HttpErrorType::ResponseDecodeFailed,
683                reason: err.to_string(),
684                status_code: Some(status.as_u16()),
685                response_body: Some(response_text),
686            },
687            format!("Failed to decode JSON response: {}", err),
688            None,
689        )
690    })
691}
692
693#[cfg(test)]
694mod tests {
695    use super::*;
696    use actix_web::ResponseError;
697    use actix_web::http::StatusCode;
698    use actix_web::http::header::{HeaderName, HeaderValue};
699    use actix_web::test::TestRequest;
700    use base64::Engine;
701    use base64::engine::general_purpose::URL_SAFE_NO_PAD;
702    use secrecy::SecretString;
703    use serde_json::json;
704
705    fn other_key() -> JwtKey {
706        JwtKey::new(&SecretString::new(
707            "a-completely-different-jwt-secret-0123456789"
708                .to_string()
709                .into(),
710        ))
711        .expect("test key")
712    }
713
714    /// Signs an arbitrary JSON payload with the same HS256 helper the production claims use, so
715    /// tests can produce claims (expired, wrong shape, legacy) the public constructors can't.
716    fn sign_json(payload: serde_json::Value, key: &JwtKey) -> String {
717        sign_hs256_claim(&payload, key).expect("signing should succeed")
718    }
719
720    fn past_timestamp(seconds_ago: i64) -> i64 {
721        (Utc::now() - Duration::seconds(seconds_ago)).timestamp()
722    }
723
724    fn future_timestamp(seconds_ahead: i64) -> i64 {
725        (Utc::now() + Duration::seconds(seconds_ahead)).timestamp()
726    }
727
728    fn answer_file(
729        id: Uuid,
730        name: &str,
731        order_number: i32,
732        size_bytes: Option<i64>,
733    ) -> SubmittedAnswerFile {
734        SubmittedAnswerFile {
735            id,
736            name: name.to_string(),
737            mime: "application/octet-stream".to_string(),
738            size_bytes,
739            order_number,
740            url: format!("http://project-331.local/api/v0/files/tmc/{name}"),
741        }
742    }
743
744    /// A downstream grader grades by position, so the request must list the files in the order the
745    /// answer records, whatever order they were resolved in.
746    #[test]
747    fn grading_request_files_are_in_answer_order_with_a_claim_for_each_file() {
748        let key = JwtKey::test_key();
749        let first = Uuid::new_v4();
750        let second = Uuid::new_v4();
751        let answer_files = vec![
752            answer_file(second, "b.txt", 1, None),
753            answer_file(first, "a.tar.zst", 0, Some(12)),
754        ];
755
756        let files = grading_request_files(Some(&answer_files), "http://project-331.local", &key)
757            .expect("the files should be built");
758
759        assert_eq!(
760            files.iter().map(|file| file.id).collect::<Vec<_>>(),
761            vec![first, second]
762        );
763        assert_eq!(files[0].size_bytes, Some(12));
764        assert_eq!(
765            files[1].size_bytes, None,
766            "an unknown size must not become a zero"
767        );
768        for (file, id) in files.iter().zip([first, second]) {
769            let (path, query) = file
770                .download_url
771                .strip_prefix("http://project-331.local/api/v0/files/claimed/")
772                .expect("a claimed-file url")
773                .split_once('?')
774                .expect("a claim in the query string");
775            assert_eq!(path, id.to_string());
776            let token = query
777                .strip_prefix(&format!("{DOWNLOAD_CLAIM_PARAM}="))
778                .expect("the claim parameter");
779            let claim = DownloadClaim::validate(token, &key).expect("the claim should validate");
780            assert_eq!(claim.file_upload_id(), id);
781        }
782    }
783
784    /// A JSON answer names no files at all, which reaches the request builder as `None`.
785    #[test]
786    fn a_json_answer_has_no_grading_request_files() {
787        let key = JwtKey::test_key();
788
789        assert!(
790            grading_request_files(None, "http://project-331.local", &key)
791                .expect("the files should be built")
792                .is_empty()
793        );
794    }
795
796    #[test]
797    fn grading_update_claim_round_trips() {
798        let key = JwtKey::test_key();
799        let submission_id = Uuid::new_v4();
800        let token = GradingUpdateClaim::expiring_in_1_day(submission_id)
801            .sign(&key)
802            .expect("signing should succeed");
803        let claim = GradingUpdateClaim::validate(&token, &key).expect("the claim should validate");
804        assert_eq!(claim.submission_id(), submission_id);
805    }
806
807    /// An expired claim must not keep authorizing grading updates.
808    #[test]
809    fn expired_grading_update_claim_is_rejected() {
810        let key = JwtKey::test_key();
811        // Well past the default 60s leeway jsonwebtoken allows for clock skew.
812        let token = sign_json(
813            json!({
814                "submission_id": Uuid::new_v4(),
815                "exp": past_timestamp(3600),
816                "iat": past_timestamp(7200),
817            }),
818            &key,
819        );
820        let err = GradingUpdateClaim::validate(&token, &key)
821            .expect_err("an expired claim must be rejected");
822        assert_eq!(err.status_code(), StatusCode::UNPROCESSABLE_ENTITY);
823    }
824
825    /// A claim signed with some other secret must not validate: this is the only thing standing
826    /// between an unauthenticated caller and writing grading results.
827    #[test]
828    fn grading_update_claim_signed_with_another_key_is_rejected() {
829        let token = GradingUpdateClaim::expiring_in_1_day(Uuid::new_v4())
830            .sign(&other_key())
831            .expect("signing should succeed");
832        GradingUpdateClaim::validate(&token, &JwtKey::test_key())
833            .expect_err("a claim signed with a foreign key must be rejected");
834    }
835
836    /// Rewriting the payload of a validly signed claim (e.g. to point at another submission)
837    /// must invalidate the signature.
838    #[test]
839    fn tampered_grading_update_claim_is_rejected() {
840        let key = JwtKey::test_key();
841        let token = GradingUpdateClaim::expiring_in_1_day(Uuid::new_v4())
842            .sign(&key)
843            .expect("signing should succeed");
844        let mut parts = token.split('.');
845        let header = parts.next().expect("header");
846        let _original_payload = parts.next().expect("payload");
847        let signature = parts.next().expect("signature");
848        // Swap in a payload naming a different submission, keeping the original signature.
849        let forged_payload = URL_SAFE_NO_PAD.encode(
850            serde_json::to_vec(&json!({
851                "submission_id": Uuid::new_v4(),
852                "exp": future_timestamp(3600),
853                "iat": Utc::now().timestamp(),
854            }))
855            .expect("json"),
856        );
857        let tampered = format!("{header}.{forged_payload}.{signature}");
858        GradingUpdateClaim::validate(&tampered, &key)
859            .expect_err("a tampered claim must be rejected");
860    }
861
862    /// The classic JWT bypass: an unsigned token declaring `alg: none` must not be accepted.
863    #[test]
864    fn unsigned_grading_update_claim_is_rejected() {
865        let header = URL_SAFE_NO_PAD.encode(br#"{"alg":"none","typ":"JWT"}"#);
866        let payload = URL_SAFE_NO_PAD.encode(
867            serde_json::to_vec(&json!({
868                "submission_id": Uuid::new_v4(),
869                "exp": future_timestamp(3600),
870                "iat": Utc::now().timestamp(),
871            }))
872            .expect("json"),
873        );
874        let token = format!("{header}.{payload}.");
875        GradingUpdateClaim::validate(&token, &JwtKey::test_key())
876            .expect_err("an unsigned (alg=none) claim must be rejected");
877    }
878
879    /// The claim types share one signing key, so a token minted for one purpose must not be
880    /// usable for another. An upload claim carries no `submission_id`, so it must not deserialize
881    /// into a grading update claim (and vice versa).
882    #[test]
883    fn claims_do_not_cross_validate_between_types() {
884        let key = JwtKey::test_key();
885        let upload_token = UploadClaim::expiring_in_1_day("tmc")
886            .sign(&key)
887            .expect("signing should succeed");
888        GradingUpdateClaim::validate(&upload_token, &key)
889            .expect_err("an upload claim must not validate as a grading update claim");
890
891        let grading_token = GradingUpdateClaim::expiring_in_1_day(Uuid::new_v4())
892            .sign(&key)
893            .expect("signing should succeed");
894        UploadClaim::validate(&grading_token, &key)
895            .expect_err("a grading update claim must not validate as an upload claim");
896    }
897
898    /// A legacy-shaped claim, carrying `expiration_time` instead of `exp`, is rejected whether or
899    /// not that timestamp has passed. Pinned so that accepting the old shape again has to be
900    /// deliberate.
901    #[test]
902    fn legacy_grading_update_claim_shape_is_rejected() {
903        let key = JwtKey::test_key();
904        let submission_id = Uuid::new_v4();
905
906        let unexpired = sign_json(
907            json!({
908                "submission_id": submission_id,
909                "expiration_time": Utc::now() + Duration::hours(1),
910            }),
911            &key,
912        );
913        let err = GradingUpdateClaim::validate(&unexpired, &key)
914            .expect_err("a claim without `exp` must be rejected");
915        assert_eq!(err.status_code(), StatusCode::UNPROCESSABLE_ENTITY);
916
917        let expired = sign_json(
918            json!({
919                "submission_id": submission_id,
920                "expiration_time": Utc::now() - Duration::hours(1),
921            }),
922            &key,
923        );
924        GradingUpdateClaim::validate(&expired, &key)
925            .expect_err("an expired legacy claim must be rejected");
926    }
927
928    /// A claim missing `exp` entirely must never be treated as non-expiring.
929    #[test]
930    fn grading_update_claim_without_an_expiry_is_rejected() {
931        let key = JwtKey::test_key();
932        let token = sign_json(
933            json!({ "submission_id": Uuid::new_v4(), "iat": Utc::now().timestamp() }),
934            &key,
935        );
936        GradingUpdateClaim::validate(&token, &key)
937            .expect_err("a claim without an expiry must be rejected");
938    }
939
940    fn extract_grading_update_claim(
941        req: actix_web::HttpRequest,
942        mut payload: Payload,
943    ) -> Result<GradingUpdateClaim, ControllerError> {
944        GradingUpdateClaim::from_request(&req, &mut payload)
945            .now_or_never()
946            .expect("the extractor resolves immediately")
947    }
948
949    #[test]
950    fn extractor_accepts_a_valid_claim_header() {
951        let key = JwtKey::test_key();
952        let submission_id = Uuid::new_v4();
953        let token = GradingUpdateClaim::expiring_in_1_day(submission_id)
954            .sign(&key)
955            .expect("signing should succeed");
956        let (req, payload) = TestRequest::default()
957            .app_data(web::Data::new(key))
958            .insert_header((EXERCISE_SERVICE_GRADING_UPDATE_CLAIM_HEADER, token.as_str()))
959            .to_http_parts();
960        let claim = extract_grading_update_claim(req, payload).expect("should extract");
961        assert_eq!(claim.submission_id(), submission_id);
962    }
963
964    #[test]
965    fn extractor_rejects_a_missing_claim_header() {
966        let (req, payload) = TestRequest::default()
967            .app_data(web::Data::new(JwtKey::test_key()))
968            .to_http_parts();
969        let err = extract_grading_update_claim(req, payload)
970            .expect_err("a request without the claim header must be rejected");
971        assert_eq!(err.status_code(), StatusCode::UNPROCESSABLE_ENTITY);
972    }
973
974    /// A non-UTF-8 header value must be a client error, not a panic.
975    #[test]
976    fn extractor_rejects_an_invalid_utf8_claim_header() {
977        let (req, payload) = TestRequest::default()
978            .app_data(web::Data::new(JwtKey::test_key()))
979            .insert_header((
980                HeaderName::from_static(EXERCISE_SERVICE_GRADING_UPDATE_CLAIM_HEADER),
981                HeaderValue::from_bytes(&[0xff, 0xfe, 0x80]).expect("header value"),
982            ))
983            .to_http_parts();
984        let err = extract_grading_update_claim(req, payload)
985            .expect_err("a non-UTF-8 claim header must be rejected");
986        assert_eq!(err.status_code(), StatusCode::UNPROCESSABLE_ENTITY);
987    }
988
989    /// Without the signing key in app data the server cannot verify anything; that is a server
990    /// misconfiguration (500), and must never be mistaken for a valid claim.
991    #[test]
992    fn extractor_reports_a_missing_jwt_key_as_a_server_error() {
993        let token = GradingUpdateClaim::expiring_in_1_day(Uuid::new_v4())
994            .sign(&JwtKey::test_key())
995            .expect("signing should succeed");
996        let (req, payload) = TestRequest::default()
997            .insert_header((EXERCISE_SERVICE_GRADING_UPDATE_CLAIM_HEADER, token.as_str()))
998            .to_http_parts();
999        let err = extract_grading_update_claim(req, payload)
1000            .expect_err("a missing JwtKey must not yield a claim");
1001        assert_eq!(err.status_code(), StatusCode::INTERNAL_SERVER_ERROR);
1002    }
1003
1004    /// A request carrying a garbage header value must be rejected before any DB work.
1005    #[test]
1006    fn extractor_rejects_a_non_jwt_claim_header() {
1007        let (req, payload) = TestRequest::default()
1008            .app_data(web::Data::new(JwtKey::test_key()))
1009            .insert_header((EXERCISE_SERVICE_GRADING_UPDATE_CLAIM_HEADER, "not-a-jwt"))
1010            .to_http_parts();
1011        let err = extract_grading_update_claim(req, payload)
1012            .expect_err("a non-JWT claim header must be rejected");
1013        assert_eq!(err.status_code(), StatusCode::UNPROCESSABLE_ENTITY);
1014    }
1015}