Skip to main content

htsget_http/middleware/
auth.rs

1//! The htsget authorization middleware.
2//!
3
4use crate::error::{Result as HtsGetResult, WrappedHtsGetError};
5use crate::middleware::error::Error::AuthBuilderError;
6use crate::middleware::error::Result;
7use crate::{Endpoint, HtsGetError};
8use cfg_if::cfg_if;
9use headers::authorization::Bearer;
10use headers::{Authorization, Header};
11use htsget_config::config::advanced::CONTEXT_HEADER_PREFIX;
12use htsget_config::config::advanced::auth::authorization::AuthorizationSource;
13use htsget_config::config::advanced::auth::jwt::JwtKey;
14use htsget_config::config::advanced::auth::response::AuthorizationRestrictionsBuilder;
15use htsget_config::config::advanced::auth::{AuthConfig, AuthorizationRestrictions};
16use htsget_config::config::advanced::callout::Callout;
17use htsget_config::config::location::{Location, PrefixOrId};
18use htsget_config::types::{Class, Interval, Query};
19use http::{HeaderMap, HeaderName, HeaderValue};
20use jsonpath_rust::JsonPath;
21use jsonwebtoken::jwk::JwkSet;
22use jsonwebtoken::{Algorithm, DecodingKey, TokenData, Validation, decode, decode_header};
23use serde::de::DeserializeOwned;
24use serde_json::Value;
25use std::fmt::{Debug, Formatter};
26use std::str::FromStr;
27use tracing::{debug, trace};
28
29/// The authorization middleware builder.
30#[derive(Default, Debug)]
31pub struct AuthBuilder {
32  config: Option<AuthConfig>,
33}
34
35impl AuthBuilder {
36  /// Set the config.
37  pub fn with_config(mut self, config: AuthConfig) -> Self {
38    self.config = Some(config);
39    self
40  }
41
42  /// Build the auth layer, ensures that the config sets the correct parameters.
43  pub fn build(self) -> Result<Auth> {
44    let Some(mut config) = self.config else {
45      return Err(AuthBuilderError("missing config".to_string()));
46    };
47
48    let mut decoding_key = None;
49    if let Some(JwtKey::PublicKey(public_key)) = config.jwt_mut() {
50      decoding_key = Some(
51        Auth::decode_public_key(public_key)
52          .map_err(|_| AuthBuilderError("failed to decode public key".to_string()))?,
53      );
54    }
55
56    Ok(Auth {
57      config,
58      decoding_key,
59    })
60  }
61}
62
63/// The auth middleware layer.
64#[derive(Clone)]
65pub struct Auth {
66  config: AuthConfig,
67  decoding_key: Option<DecodingKey>,
68}
69
70impl Debug for Auth {
71  fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
72    f.debug_struct("config").finish()
73  }
74}
75
76const ENDPOINT_TYPE_HEADER_NAME: &str = "Endpoint-Type";
77const ID_HEADER_NAME: &str = "Id";
78
79impl Auth {
80  /// Get the config for this auth layer instance.
81  pub fn config(&self) -> &AuthConfig {
82    &self.config
83  }
84
85  /// Fetch a JSON object from the callout.
86  pub async fn fetch_from_callout<D: DeserializeOwned>(
87    callout: &mut Callout,
88    headers: HeaderMap,
89  ) -> HtsGetResult<D> {
90    let url = callout.url().to_string();
91    trace!("fetching url: {}", url);
92
93    let forwarded_header_names: Vec<String> =
94      headers.keys().map(|k| k.as_str().to_string()).collect();
95
96    // Build the client with forwarded headers for identity based cache.
97    let http = callout.http_mut();
98    let client = http
99      .as_inner_built_with_forwarded_headers(&forwarded_header_names)
100      .map_err(|err| HtsGetError::InternalError(format!("failed to fetch data from {url}: {err}")))?
101      .clone();
102
103    let ttl_ceiling_secs = callout.http().ttl_ceiling_secs();
104    let mut request_headers = headers;
105    if let Some(ceiling_secs) = ttl_ceiling_secs {
106      let cache_control_value = format!("max-age={ceiling_secs}");
107      request_headers.insert(
108        http::header::CACHE_CONTROL,
109        cache_control_value
110          .parse()
111          .expect("valid cache-control header value"),
112      );
113    }
114
115    let response = client.get(&url).headers(request_headers).send().await?;
116    trace!("response: {:?}", response);
117
118    let status = response.status();
119
120    // Forward a valid htsget error if that's what the backend returns.
121    let value = response.json::<Value>().await.map_err(|err| {
122      HtsGetError::InternalError(format!("failed to fetch data from {url}: {err}"))
123    })?;
124    trace!("value: {}", value);
125
126    match serde_json::from_value::<D>(value.clone()) {
127      Ok(response) => Ok(response),
128      Err(_) => match serde_json::from_value::<WrappedHtsGetError>(value.clone()) {
129        Ok(err) => Err(HtsGetError::Wrapped(err, status)),
130        Err(_) => Err(HtsGetError::InternalError(format!(
131          "failed to fetch data from {url}: {value}"
132        ))),
133      },
134    }
135  }
136
137  /// Get a decoding key from the JWKS callout.
138  pub async fn decode_jwks(callout: &mut Callout, token: &str) -> HtsGetResult<DecodingKey> {
139    // Decode header and get the key id.
140    let header = decode_header(token)?;
141    let kid = header
142      .kid
143      .ok_or_else(|| HtsGetError::PermissionDenied("JWT missing key ID".to_string()))?;
144
145    // Fetch JWKS from the authorization server and find matching JWK.
146    let jwks = Self::fetch_from_callout::<JwkSet>(callout, Default::default()).await?;
147    let matched_jwk = jwks
148      .find(&kid)
149      .ok_or_else(|| HtsGetError::PermissionDenied("matching JWK not found".to_string()))?;
150
151    Ok(DecodingKey::from_jwk(matched_jwk)?)
152  }
153
154  /// Decode a public key into an RSA, EdDSA or ECDSA pem-formatted decoding key.
155  pub fn decode_public_key(key: &[u8]) -> HtsGetResult<DecodingKey> {
156    Ok(
157      DecodingKey::from_rsa_pem(key)
158        .or_else(|_| DecodingKey::from_ed_pem(key))
159        .or_else(|_| DecodingKey::from_ec_pem(key))?,
160    )
161  }
162
163  /// Build the headers to send, allow-listed client headers are forwarded
164  /// directly and context values take `Htsget-Context-` as a prefix.
165  pub fn forwarded_headers(
166    callout: &Callout,
167    request_headers: &HeaderMap,
168    request_extensions: Option<Value>,
169    request_endpoint: &Endpoint,
170    id: &str,
171  ) -> HtsGetResult<HeaderMap> {
172    let forward = callout.forward();
173    let mut forwarded_headers = forward.headers().filter(request_headers);
174
175    let context = forward.context();
176
177    if let Some(request_extensions) = request_extensions {
178      for extension in context.extensions() {
179        let Some(value) = request_extensions.query(extension.json_path()).ok() else {
180          continue;
181        };
182
183        let value = value.first().ok_or_else(|| {
184          HtsGetError::InternalError("extension does not have only one value".to_string())
185        })?;
186        let value = value.as_str().ok_or_else(|| {
187          HtsGetError::InternalError("extension value is not a string".to_string())
188        })?;
189
190        let header_name =
191          HeaderName::from_str(&format!("{}{}", CONTEXT_HEADER_PREFIX, extension.name()))?;
192        let value = HeaderValue::from_str(value)?;
193        forwarded_headers.insert(header_name, value);
194      }
195    }
196
197    if context.endpoint_type() {
198      let header_name = HeaderName::from_str(&format!(
199        "{}{}",
200        CONTEXT_HEADER_PREFIX, ENDPOINT_TYPE_HEADER_NAME
201      ))?;
202      let value = HeaderValue::from_str(&request_endpoint.to_string())?;
203
204      forwarded_headers.insert(header_name, value);
205    }
206
207    if context.id() {
208      let header_name =
209        HeaderName::from_str(&format!("{}{}", CONTEXT_HEADER_PREFIX, ID_HEADER_NAME))?;
210      let value = HeaderValue::from_str(id)?;
211
212      forwarded_headers.insert(header_name, value);
213    }
214
215    Ok(forwarded_headers)
216  }
217
218  /// Query the authorization source to get the restrictions. The claims are assumed to
219  /// be valid.
220  pub async fn query_authorization_service(
221    &mut self,
222    headers: &HeaderMap,
223    request_extensions: Option<Value>,
224    request_endpoint: &Endpoint,
225    id: &str,
226  ) -> HtsGetResult<Option<AuthorizationRestrictions>> {
227    match self.config.authorization_mut() {
228      Some(AuthorizationSource::Callout(callout)) => {
229        let forwarded_headers =
230          Self::forwarded_headers(callout, headers, request_extensions, request_endpoint, id)?;
231
232        Self::fetch_from_callout(callout, forwarded_headers)
233          .await
234          .map(Some)
235      }
236      Some(AuthorizationSource::Static(restrictions)) => Ok(Some(restrictions.clone())),
237      None => Ok(None),
238    }
239  }
240
241  /// Validate the restrictions, returning an error if the user is not authorized.
242  /// If `suppressed_interval` is set then no error is returning if there is a
243  /// path match but no restrictions match. Instead, as many regions as possible
244  /// are returned.
245  pub fn validate_restrictions(
246    restrictions: AuthorizationRestrictions,
247    path: &str,
248    queries: &mut [Query],
249    suppressed_interval: bool,
250  ) -> HtsGetResult<AuthorizationRestrictions> {
251    // Find all rules matching the path.
252    let matching_rules = restrictions
253      .into_rules()
254      .into_iter()
255      .filter(|rule| {
256        match rule.location() {
257          Location::Simple(location) if location.prefix_or_id().is_some() => {
258            match location.prefix_or_id().unwrap_or_default() {
259              PrefixOrId::Prefix(prefix) => {
260                // A prefix has a starts with rule.
261                path.starts_with(&prefix)
262              }
263              PrefixOrId::Id(id) => {
264                // An id location must match exactly.
265                id == path
266              }
267            }
268          }
269          Location::Regex(location) => {
270            // A regex location matches using the regex.
271            location.regex().is_match(path)
272          }
273          // Missing valid location.
274          _ => false,
275        }
276      })
277      .collect::<Vec<_>>();
278
279    // If no paths match, then this is an authorization error.
280    if matching_rules.is_empty() {
281      return Err(HtsGetError::PermissionDenied(
282        "failed to authorize user based on authorization service restrictions".to_string(),
283      ));
284    }
285
286    let (allows_all, allows_specific): (Vec<_>, Vec<_>) = matching_rules
287      .into_iter()
288      .partition(|rule| rule.rules().is_none());
289
290    // Otherwise, we need to check if the specific reference name is allowed for all queries.
291    for query in queries {
292      // If the request is for headers only, then this should always be allowed.
293      if query.class() == Class::Header {
294        continue;
295      }
296
297      let matching_restriction = allows_specific
298        .iter()
299        .flat_map(|rule| rule.rules().unwrap_or_default())
300        .filter_map(|restriction| {
301          // The reference name should match exactly if it's set, otherwise allow any reference name.
302          let name_match = restriction.reference_name().is_none()
303            || restriction.reference_name() == query.reference_name();
304          // The format should match if it's defined, otherwise it allows any format.
305          let format_match =
306            restriction.format().is_none() || restriction.format() == Some(query.format());
307          // The interval should match and be constrained, considering undefined start or end ranges.
308          let interval_match = if suppressed_interval {
309            restriction.interval().constraint_interval(query.interval())
310          } else {
311            restriction.interval().contains_interval(query.interval())
312          };
313
314          if let Some(interval_match) = interval_match
315            && name_match
316            && format_match
317          {
318            return Some(interval_match);
319          }
320
321          None
322        })
323        .max_by(Interval::order_by_range); // The largest interval should be used if there are multiple matches.
324
325      if suppressed_interval {
326        if allows_all.is_empty() && matching_restriction.is_none() {
327          // If nothing allows all and there are no matching intervals then return an empty response.
328          query.set_class(Class::Header);
329          continue;
330        }
331
332        if let Some(matching_restriction) = matching_restriction {
333          query.set_interval(matching_restriction);
334        }
335      } else if allows_all.is_empty() && matching_restriction.is_none() {
336        return Err(HtsGetError::PermissionDenied(
337          "failed to authorize user based on authorization service restrictions".to_string(),
338        ));
339      }
340    }
341
342    AuthorizationRestrictionsBuilder::default()
343      .rules([allows_all, allows_specific].concat())
344      .build()
345      .map_err(|err| HtsGetError::InternalError(err.to_string()))
346  }
347
348  /// Validate only the JWT without looking up restrictions and validating those. Returns the
349  /// decoded JWT token.
350  pub async fn validate_jwt(&mut self, headers: &HeaderMap) -> HtsGetResult<TokenData<Value>> {
351    let auth_token = headers
352      .values()
353      .find_map(|value| Authorization::<Bearer>::decode(&mut [value].into_iter()).ok())
354      .ok_or_else(|| {
355        HtsGetError::InvalidAuthentication("invalid authorization header".to_string())
356      })?;
357
358    let owned_jwks_key;
359    let decoding_key = if let Some(ref decoding_key) = self.decoding_key {
360      decoding_key
361    } else if let Some(JwtKey::Jwks(callout)) = self.config.jwt_mut() {
362      owned_jwks_key = Self::decode_jwks(callout, auth_token.token()).await?;
363      &owned_jwks_key
364    } else {
365      return Err(HtsGetError::InternalError(
366        "JWT validation not set".to_string(),
367      ));
368    };
369
370    // Decode and validate the JWT
371    let mut validation = Validation::default();
372    validation.validate_exp = true;
373    validation.validate_aud = true;
374    validation.validate_nbf = true;
375
376    if let Some(iss) = self.config.validate_issuer() {
377      validation.set_issuer(iss);
378      validation.required_spec_claims.insert("iss".to_string());
379    }
380    if let Some(aud) = self.config.validate_audience() {
381      validation.set_audience(aud);
382      validation.required_spec_claims.insert("aud".to_string());
383    }
384    if let Some(sub) = self.config.validate_subject() {
385      validation.sub = Some(sub.to_string());
386      validation.required_spec_claims.insert("sub".to_string());
387    }
388
389    // Each supported algorithm must be tried individually because the jsonwebtoken validation
390    // logic only tries one algorithm in the vec: https://github.com/Keats/jsonwebtoken/issues/297
391    validation.algorithms = vec![Algorithm::RS256];
392    let decoded_claims = decode::<Value>(auth_token.token(), decoding_key, &validation)
393      .or_else(|_| {
394        validation.algorithms = vec![Algorithm::ES256];
395        decode::<Value>(auth_token.token(), decoding_key, &validation)
396      })
397      .or_else(|_| {
398        validation.algorithms = vec![Algorithm::EdDSA];
399        decode::<Value>(auth_token.token(), decoding_key, &validation)
400      });
401
402    let claims = match decoded_claims {
403      Ok(claims) => claims,
404      Err(err) => return Err(HtsGetError::PermissionDenied(format!("invalid JWT: {err}"))),
405    };
406
407    Ok(claims)
408  }
409
410  /// Validate the authorization flow, returning an error if the user is not authorized.
411  /// This performs the following steps:
412  ///
413  /// 1. Finds the JWT decoding key from the config or by querying a JWKS url.
414  /// 2. Validates the JWT token according to the config.
415  /// 3. Queries the authorization service for restrictions based on the config or JWT claims.
416  /// 4. Validates the restrictions to determine if the user is authorized.
417  pub async fn validate_authorization(
418    &mut self,
419    headers: &HeaderMap,
420    path: &str,
421    queries: &mut [Query],
422    request_extensions: Option<Value>,
423    endpoint: &Endpoint,
424  ) -> HtsGetResult<Option<AuthorizationRestrictions>> {
425    let restrictions = self
426      .query_authorization_service(headers, request_extensions, endpoint, path)
427      .await?;
428
429    debug!(restrictions = ?restrictions, "restrictions");
430
431    if let Some(restrictions) = restrictions {
432      cfg_if! {
433        if #[cfg(feature = "experimental")] {
434          Self::validate_restrictions(restrictions, path, queries, self.config.suppress_errors()).map(Some)
435        } else {
436          Self::validate_restrictions(restrictions, path, queries, false).map(Some)
437        }
438      }
439    } else {
440      Ok(None)
441    }
442  }
443}
444
445#[cfg(test)]
446mod tests {
447  use super::*;
448  use crate::{Endpoint, convert_to_query, match_format_from_query};
449  use htsget_config::config::advanced::HttpClient;
450  use htsget_config::config::advanced::auth::AuthConfigBuilder;
451  use htsget_config::config::advanced::auth::authorization::AuthorizationSourceBuilder;
452  use htsget_config::config::advanced::auth::response::{
453    AuthorizationRestrictionsBuilder, AuthorizationRuleBuilder, ReferenceNameRestrictionBuilder,
454  };
455  use htsget_config::config::advanced::callout::{
456    ContextExtension, ContextRules, Forward, HeaderRules,
457  };
458  use htsget_config::config::advanced::regex_location::RegexLocation;
459  use htsget_config::config::location::SimpleLocation;
460  use htsget_config::http::client::HttpClientConfig;
461  use htsget_config::types::{Format, Request};
462  use htsget_test::util::generate_key_pair;
463  use http::{HeaderMap, Uri};
464  use regex::Regex;
465  use serde_json::json;
466  use std::collections::HashMap;
467
468  #[test]
469  fn auth_builder_missing_config() {
470    let result = AuthBuilder::default().build();
471    assert!(matches!(result, Err(AuthBuilderError(_))));
472  }
473
474  #[test]
475  fn auth_builder_success_with_public_key() {
476    let (_, public_key) = generate_key_pair();
477
478    let config = create_test_auth_config(public_key);
479    let result = AuthBuilder::default().with_config(config).build();
480    assert!(result.is_ok());
481  }
482
483  #[test]
484  fn validate_restrictions_rule_allows_all() {
485    let rule = AuthorizationRuleBuilder::default()
486      .location(test_location())
487      .build()
488      .unwrap();
489    let restrictions = AuthorizationRestrictionsBuilder::default()
490      .rule(rule)
491      .build()
492      .unwrap();
493
494    let request = create_test_query(Endpoint::Reads, "sample1", HashMap::new());
495    let result =
496      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], false);
497    assert!(result.is_ok());
498  }
499
500  #[test]
501  fn validate_restrictions_exact_path_match() {
502    let reference_restriction = ReferenceNameRestrictionBuilder::default()
503      .name("chr1")
504      .format(Format::Bam)
505      .start(1000)
506      .end(2000)
507      .build()
508      .unwrap();
509    let rule = AuthorizationRuleBuilder::default()
510      .location(test_location())
511      .reference_name(reference_restriction)
512      .build()
513      .unwrap();
514    let restrictions = AuthorizationRestrictionsBuilder::default()
515      .rule(rule)
516      .build()
517      .unwrap();
518
519    let mut query = HashMap::new();
520    query.insert("referenceName".to_string(), "chr1".to_string());
521    query.insert("start".to_string(), "1500".to_string());
522    query.insert("end".to_string(), "1800".to_string());
523    query.insert("format".to_string(), "BAM".to_string());
524
525    let request = create_test_query(Endpoint::Reads, "sample1", query);
526    let result =
527      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], false);
528    assert!(result.is_ok());
529  }
530
531  #[test]
532  fn validate_restrictions_regex_prefix_match() {
533    let reference_restriction = ReferenceNameRestrictionBuilder::default()
534      .name("chr1")
535      .format(Format::Bam)
536      .build()
537      .unwrap();
538    let rule = AuthorizationRuleBuilder::default()
539      .location(Location::Simple(Box::new(SimpleLocation::new(
540        Default::default(),
541        "".to_string(),
542        Some(PrefixOrId::Prefix("sam".to_string())),
543      ))))
544      .reference_name(reference_restriction)
545      .build()
546      .unwrap();
547    let restrictions = AuthorizationRestrictionsBuilder::default()
548      .rule(rule)
549      .build()
550      .unwrap();
551
552    let mut query = HashMap::new();
553    query.insert("referenceName".to_string(), "chr1".to_string());
554    query.insert("format".to_string(), "BAM".to_string());
555
556    let request = create_test_query(Endpoint::Reads, "sample123", query);
557    let result =
558      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], false);
559    assert!(result.is_ok());
560  }
561
562  #[test]
563  fn validate_restrictions_regex_match() {
564    let reference_restriction = ReferenceNameRestrictionBuilder::default()
565      .name("chr1")
566      .format(Format::Bam)
567      .build()
568      .unwrap();
569    let rule = AuthorizationRuleBuilder::default()
570      .location(Location::Regex(Box::new(RegexLocation::new(
571        Regex::new("sample(.+)").unwrap(),
572        "".to_string(),
573        Default::default(),
574      ))))
575      .reference_name(reference_restriction)
576      .build()
577      .unwrap();
578    let restrictions = AuthorizationRestrictionsBuilder::default()
579      .rule(rule)
580      .build()
581      .unwrap();
582
583    let mut query = HashMap::new();
584    query.insert("referenceName".to_string(), "chr1".to_string());
585    query.insert("format".to_string(), "BAM".to_string());
586
587    let request = create_test_query(Endpoint::Reads, "sample123", query);
588    let result =
589      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], false);
590    assert!(result.is_ok());
591  }
592
593  #[test]
594  fn validate_restrictions_forward_headers() {
595    let request_headers = HeaderMap::from_iter([
596      (
597        "Authorization".parse::<http::HeaderName>().unwrap(),
598        "Bearer Value".parse().unwrap(),
599      ),
600      ("Custom1".parse().unwrap(), "Value".parse().unwrap()),
601      ("Custom2".parse().unwrap(), "Value".parse().unwrap()),
602    ]);
603
604    // Forward Authorization and Custom1
605    let callout = callout_with_forward(Forward::new(
606      HeaderRules::new(
607        vec!["Authorization".to_string(), "Custom1".to_string()],
608        vec![],
609      ),
610      ContextRules::default(),
611    ));
612    let forwarded_headers =
613      Auth::forwarded_headers(&callout, &request_headers, None, &Endpoint::Reads, "id").unwrap();
614    assert_eq!(
615      forwarded_headers,
616      HeaderMap::from_iter([
617        (
618          "Authorization".parse::<http::HeaderName>().unwrap(),
619          "Bearer Value".parse().unwrap()
620        ),
621        ("Custom1".parse().unwrap(), "Value".parse().unwrap()),
622      ])
623    );
624
625    // Wildcard allow with specific deny.
626    let callout = callout_with_forward(Forward::new(
627      HeaderRules::new(vec!["Custom*".to_string()], vec!["Custom2".to_string()]),
628      ContextRules::default(),
629    ));
630    let forwarded_headers =
631      Auth::forwarded_headers(&callout, &request_headers, None, &Endpoint::Reads, "id").unwrap();
632    assert_eq!(
633      forwarded_headers,
634      HeaderMap::from_iter([(
635        "Custom1".parse::<http::HeaderName>().unwrap(),
636        "Value".parse().unwrap()
637      ),])
638    );
639
640    // Extension is inserted with Htsget-Context- prefix.
641    let callout = callout_with_forward(Forward::new(
642      HeaderRules::default(),
643      ContextRules::new(
644        false,
645        false,
646        vec![ContextExtension::new(
647          "$.Key".to_string(),
648          "Custom1".to_string(),
649        )],
650      ),
651    ));
652    let forwarded_headers = Auth::forwarded_headers(
653      &callout,
654      &request_headers,
655      Some(json!({ "Key": "Value" })),
656      &Endpoint::Reads,
657      "id",
658    )
659    .unwrap();
660    assert_eq!(
661      forwarded_headers,
662      HeaderMap::from_iter([(
663        format!("{}Custom1", CONTEXT_HEADER_PREFIX).parse().unwrap(),
664        "Value".parse().unwrap()
665      ),])
666    );
667
668    // endpoint_type context header.
669    let callout = callout_with_forward(Forward::new(
670      HeaderRules::default(),
671      ContextRules::new(true, false, vec![]),
672    ));
673    let forwarded_headers =
674      Auth::forwarded_headers(&callout, &request_headers, None, &Endpoint::Variants, "id").unwrap();
675    assert_eq!(
676      forwarded_headers,
677      HeaderMap::from_iter([(
678        format!("{}{}", CONTEXT_HEADER_PREFIX, ENDPOINT_TYPE_HEADER_NAME)
679          .parse()
680          .unwrap(),
681        "variants".parse().unwrap()
682      ),])
683    );
684
685    // id context header.
686    let callout = callout_with_forward(Forward::new(
687      HeaderRules::default(),
688      ContextRules::new(false, true, vec![]),
689    ));
690    let forwarded_headers =
691      Auth::forwarded_headers(&callout, &request_headers, None, &Endpoint::Variants, "id").unwrap();
692    assert_eq!(
693      forwarded_headers,
694      HeaderMap::from_iter([(
695        format!("{}{}", CONTEXT_HEADER_PREFIX, ID_HEADER_NAME)
696          .parse()
697          .unwrap(),
698        "id".parse().unwrap()
699      ),])
700    );
701  }
702
703  #[test]
704  fn validate_restrictions_reference_name_mismatch() {
705    let reference_restriction = ReferenceNameRestrictionBuilder::default()
706      .name("chr1")
707      .format(Format::Bam)
708      .build()
709      .unwrap();
710    let rule = AuthorizationRuleBuilder::default()
711      .location(test_location())
712      .reference_name(reference_restriction)
713      .build()
714      .unwrap();
715    let restrictions = AuthorizationRestrictionsBuilder::default()
716      .rule(rule.clone())
717      .build()
718      .unwrap();
719
720    let mut query = HashMap::new();
721    query.insert("class".to_string(), "header".to_string());
722    query.insert("format".to_string(), "BAM".to_string());
723
724    let request = create_test_query(Endpoint::Reads, "sample1", query);
725    let result =
726      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], false);
727    assert!(result.is_ok());
728  }
729
730  #[test]
731  fn validate_restrictions_header() {
732    let reference_restriction = ReferenceNameRestrictionBuilder::default()
733      .name("chr1")
734      .format(Format::Bam)
735      .build()
736      .unwrap();
737    let rule = AuthorizationRuleBuilder::default()
738      .location(test_location())
739      .reference_name(reference_restriction)
740      .build()
741      .unwrap();
742    let restrictions = AuthorizationRestrictionsBuilder::default()
743      .rule(rule.clone())
744      .build()
745      .unwrap();
746
747    let mut query = HashMap::new();
748    query.insert("format".to_string(), "BAM".to_string());
749    query.insert("class".to_string(), "header".to_string());
750
751    let request = create_test_query(Endpoint::Reads, "sample1", query);
752    let result =
753      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], false);
754    assert!(result.is_ok());
755  }
756
757  #[cfg(feature = "experimental")]
758  #[test]
759  fn validate_restrictions_reference_name_mismatch_suppressed() {
760    let reference_restriction = ReferenceNameRestrictionBuilder::default()
761      .name("chr1")
762      .format(Format::Bam)
763      .build()
764      .unwrap();
765    let rule = AuthorizationRuleBuilder::default()
766      .location(test_location())
767      .reference_name(reference_restriction)
768      .build()
769      .unwrap();
770    let restrictions = AuthorizationRestrictionsBuilder::default()
771      .rule(rule.clone())
772      .build()
773      .unwrap();
774
775    let mut query = HashMap::new();
776    query.insert("referenceName".to_string(), "chr2".to_string());
777    query.insert("format".to_string(), "BAM".to_string());
778
779    let request = create_test_query(Endpoint::Reads, "sample1", query);
780    let result =
781      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], true);
782    assert!(result.is_ok());
783  }
784
785  #[test]
786  fn validate_restrictions_format_mismatch() {
787    let reference_restriction = ReferenceNameRestrictionBuilder::default()
788      .name("chr1")
789      .format(Format::Bam)
790      .build()
791      .unwrap();
792    let rule = AuthorizationRuleBuilder::default()
793      .location(test_location())
794      .reference_name(reference_restriction)
795      .build()
796      .unwrap();
797    let restrictions = AuthorizationRestrictionsBuilder::default()
798      .rule(rule.clone())
799      .build()
800      .unwrap();
801
802    let mut query = HashMap::new();
803    query.insert("referenceName".to_string(), "chr1".to_string());
804    query.insert("format".to_string(), "CRAM".to_string());
805
806    let request = create_test_query(Endpoint::Reads, "sample1", query);
807    let result =
808      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], false);
809    assert!(result.is_err());
810  }
811
812  #[cfg(feature = "experimental")]
813  #[test]
814  fn validate_restrictions_format_mismatch_suppressed() {
815    let reference_restriction = ReferenceNameRestrictionBuilder::default()
816      .name("chr1")
817      .format(Format::Bam)
818      .build()
819      .unwrap();
820    let rule = AuthorizationRuleBuilder::default()
821      .location(test_location())
822      .reference_name(reference_restriction)
823      .build()
824      .unwrap();
825    let restrictions = AuthorizationRestrictionsBuilder::default()
826      .rule(rule.clone())
827      .build()
828      .unwrap();
829
830    let mut query = HashMap::new();
831    query.insert("referenceName".to_string(), "chr1".to_string());
832    query.insert("format".to_string(), "CRAM".to_string());
833
834    let request = create_test_query(Endpoint::Reads, "sample1", query);
835    let result =
836      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], true);
837    assert!(result.is_ok());
838  }
839
840  #[test]
841  fn validate_restrictions_interval_not_contained() {
842    // Restriction:       1000----------2000
843    // Request:               1250--1750
844    // Result:                1250--1750
845    test_interval_suppressed(
846      Some(1000),
847      Some(2000),
848      Some(1250),
849      Some(1750),
850      (Interval::new(Some(1250), Some(1750)), Class::Body),
851      false,
852      false,
853    );
854
855    // Restriction:       1000----------2000
856    // Request:   500------------------------------->
857    // Result:                   err
858    test_interval_suppressed(
859      Some(1000),
860      Some(2000),
861      Some(500),
862      None,
863      (Interval::new(Some(500), None), Class::Body),
864      true,
865      false,
866    );
867
868    // Restriction:       1000----------2000
869    // Request:   <------------------------------2500
870    // Result:                   err
871    test_interval_suppressed(
872      Some(1000),
873      Some(2000),
874      None,
875      Some(2500),
876      (Interval::new(None, Some(2500)), Class::Body),
877      true,
878      false,
879    );
880
881    // Restriction:       1000----------2000
882    // Request:   <--------------------------------->
883    // Result:                   err
884    test_interval_suppressed(
885      Some(1000),
886      Some(2000),
887      None,
888      None,
889      (Interval::new(None, None), Class::Body),
890      true,
891      false,
892    );
893
894    // Restriction:       1000----------2000
895    // Request:   500------------1500
896    // Result:                   err
897    test_interval_suppressed(
898      Some(1000),
899      Some(2000),
900      Some(500),
901      Some(1500),
902      (Interval::new(Some(500), Some(1500)), Class::Body),
903      true,
904      false,
905    );
906
907    // Restriction:       1000----------2000
908    // Request:   <--------------1500
909    // Result:                   err
910    test_interval_suppressed(
911      Some(1000),
912      Some(2000),
913      None,
914      Some(1500),
915      (Interval::new(None, Some(1500)), Class::Body),
916      true,
917      false,
918    );
919
920    // Restriction:       1000----------2000
921    // Request:                  1500------------2500
922    // Result:                   err
923    test_interval_suppressed(
924      Some(1000),
925      Some(2000),
926      Some(1500),
927      Some(2500),
928      (Interval::new(Some(1500), Some(2500)), Class::Body),
929      true,
930      false,
931    );
932
933    // Restriction:       1000----------2000
934    // Request:                  1500--------------->
935    // Result:                   err
936    test_interval_suppressed(
937      Some(1000),
938      Some(2000),
939      Some(1500),
940      None,
941      (Interval::new(Some(1500), None), Class::Body),
942      true,
943      false,
944    );
945
946    // Restriction:       1000----------2000
947    // Request:   500-----1000
948    // Result:                   err
949    test_interval_suppressed(
950      Some(1000),
951      Some(2000),
952      Some(500),
953      Some(1000),
954      (Interval::new(Some(500), Some(1000)), Class::Body),
955      true,
956      false,
957    );
958
959    // Restriction:       1000----------2000
960    // Request:                         2000-----2500
961    // Result:                   err
962    test_interval_suppressed(
963      Some(1000),
964      Some(2000),
965      Some(2000),
966      Some(2500),
967      (Interval::new(Some(2000), Some(2500)), Class::Body),
968      true,
969      false,
970    );
971
972    // Restriction:       <-------------2000
973    // Request:   500------------1500
974    // Result:    500------------1500
975    test_interval_suppressed(
976      None,
977      Some(2000),
978      Some(500),
979      Some(1500),
980      (Interval::new(Some(500), Some(1500)), Class::Body),
981      false,
982      false,
983    );
984
985    // Restriction:       <-------------2000
986    // Request:                  1500------------2500
987    // Result:                   err
988    test_interval_suppressed(
989      None,
990      Some(2000),
991      Some(1500),
992      Some(2500),
993      (Interval::new(Some(1500), Some(2500)), Class::Body),
994      true,
995      false,
996    );
997
998    // Restriction:       1000------------->
999    // Request:                  1500------------2500
1000    // Result:                   1500------------2500
1001    test_interval_suppressed(
1002      Some(1000),
1003      None,
1004      Some(1500),
1005      Some(2500),
1006      (Interval::new(Some(1500), Some(2500)), Class::Body),
1007      false,
1008      false,
1009    );
1010
1011    // Restriction:       1000------------->
1012    // Request:   500------------1500
1013    // Result:                   err
1014    test_interval_suppressed(
1015      Some(1000),
1016      None,
1017      Some(500),
1018      Some(1500),
1019      (Interval::new(Some(500), Some(1500)), Class::Body),
1020      true,
1021      false,
1022    );
1023
1024    // Restriction:       <---------------->
1025    // Request:   500----------------------------2500
1026    // Result:    500----------------------------2500
1027    test_interval_suppressed(
1028      None,
1029      None,
1030      Some(500),
1031      Some(2500),
1032      (Interval::new(Some(500), Some(2500)), Class::Body),
1033      false,
1034      false,
1035    );
1036
1037    // Restriction:       <---------------->
1038    // Request:   500------------------------------->
1039    // Result:    500------------------------------->
1040    test_interval_suppressed(
1041      None,
1042      None,
1043      Some(500),
1044      None,
1045      (Interval::new(Some(500), None), Class::Body),
1046      false,
1047      false,
1048    );
1049
1050    // Restriction:       <---------------->
1051    // Request:   <------------------------------2500
1052    // Result:    <------------------------------2500
1053    test_interval_suppressed(
1054      None,
1055      None,
1056      None,
1057      Some(2500),
1058      (Interval::new(None, Some(2500)), Class::Body),
1059      false,
1060      false,
1061    );
1062  }
1063
1064  #[cfg(feature = "experimental")]
1065  #[test]
1066  fn validate_restrictions_interval_suppressed() {
1067    // Restriction:       1000----------2000
1068    // Request:               1250--1750
1069    // Result:                1250--1750
1070    test_interval_suppressed(
1071      Some(1000),
1072      Some(2000),
1073      Some(1250),
1074      Some(1750),
1075      (Interval::new(Some(1250), Some(1750)), Class::Body),
1076      false,
1077      true,
1078    );
1079
1080    // Restriction:       1000----------2000
1081    // Request:   500------------------------------->
1082    // Result:            1000----------2000
1083    test_interval_suppressed(
1084      Some(1000),
1085      Some(2000),
1086      Some(500),
1087      None,
1088      (Interval::new(Some(1000), Some(2000)), Class::Body),
1089      false,
1090      true,
1091    );
1092
1093    // Restriction:       1000----------2000
1094    // Request:   <------------------------------2500
1095    // Result:            1000----------2000
1096    test_interval_suppressed(
1097      Some(1000),
1098      Some(2000),
1099      None,
1100      Some(2500),
1101      (Interval::new(Some(1000), Some(2000)), Class::Body),
1102      false,
1103      true,
1104    );
1105
1106    // Restriction:       1000----------2000
1107    // Request:   <--------------------------------->
1108    // Result:            1000----------2000
1109    test_interval_suppressed(
1110      Some(1000),
1111      Some(2000),
1112      None,
1113      None,
1114      (Interval::new(Some(1000), Some(2000)), Class::Body),
1115      false,
1116      true,
1117    );
1118
1119    // Restriction:       1000----------2000
1120    // Request:   500------------1500
1121    // Result:            1000---1500
1122    test_interval_suppressed(
1123      Some(1000),
1124      Some(2000),
1125      Some(500),
1126      Some(1500),
1127      (Interval::new(Some(1000), Some(1500)), Class::Body),
1128      false,
1129      true,
1130    );
1131
1132    // Restriction:       1000----------2000
1133    // Request:   <--------------1500
1134    // Result:            1000---1500
1135    test_interval_suppressed(
1136      Some(1000),
1137      Some(2000),
1138      None,
1139      Some(1500),
1140      (Interval::new(Some(1000), Some(1500)), Class::Body),
1141      false,
1142      true,
1143    );
1144
1145    // Restriction:       1000----------2000
1146    // Request:                  1500------------2500
1147    // Result:                   1500---2000
1148    test_interval_suppressed(
1149      Some(1000),
1150      Some(2000),
1151      Some(1500),
1152      Some(2500),
1153      (Interval::new(Some(1500), Some(2000)), Class::Body),
1154      false,
1155      true,
1156    );
1157
1158    // Restriction:       1000----------2000
1159    // Request:                  1500--------------->
1160    // Result:                   1500---2000
1161    test_interval_suppressed(
1162      Some(1000),
1163      Some(2000),
1164      Some(1500),
1165      None,
1166      (Interval::new(Some(1500), Some(2000)), Class::Body),
1167      false,
1168      true,
1169    );
1170
1171    // Restriction:       1000----------2000
1172    // Request:   500-----1000
1173    // Result:            -
1174    test_interval_suppressed(
1175      Some(1000),
1176      Some(2000),
1177      Some(500),
1178      Some(1000),
1179      (Interval::new(Some(500), Some(1000)), Class::Header),
1180      false,
1181      true,
1182    );
1183
1184    // Restriction:       1000----------2000
1185    // Request:                         2000-----2500
1186    // Result:                          -
1187    test_interval_suppressed(
1188      Some(1000),
1189      Some(2000),
1190      Some(2000),
1191      Some(2500),
1192      (Interval::new(Some(2000), Some(2500)), Class::Header),
1193      false,
1194      true,
1195    );
1196
1197    // Restriction:       <-------------2000
1198    // Request:   500------------1500
1199    // Result:    500------------1500
1200    test_interval_suppressed(
1201      None,
1202      Some(2000),
1203      Some(500),
1204      Some(1500),
1205      (Interval::new(Some(500), Some(1500)), Class::Body),
1206      false,
1207      true,
1208    );
1209
1210    // Restriction:       <-------------2000
1211    // Request:                  1500------------2500
1212    // Result:                   1500---2000
1213    test_interval_suppressed(
1214      None,
1215      Some(2000),
1216      Some(1500),
1217      Some(2500),
1218      (Interval::new(Some(1500), Some(2000)), Class::Body),
1219      false,
1220      true,
1221    );
1222
1223    // Restriction:       1000------------->
1224    // Request:                  1500------------2500
1225    // Result:                   1500------------2500
1226    test_interval_suppressed(
1227      Some(1000),
1228      None,
1229      Some(1500),
1230      Some(2500),
1231      (Interval::new(Some(1500), Some(2500)), Class::Body),
1232      false,
1233      true,
1234    );
1235
1236    // Restriction:       1000------------->
1237    // Request:   500------------1500
1238    // Result:            1000---1500
1239    test_interval_suppressed(
1240      Some(1000),
1241      None,
1242      Some(500),
1243      Some(1500),
1244      (Interval::new(Some(1000), Some(1500)), Class::Body),
1245      false,
1246      true,
1247    );
1248
1249    // Restriction:       <---------------->
1250    // Request:   500----------------------------2500
1251    // Result:    500----------------------------2500
1252    test_interval_suppressed(
1253      None,
1254      None,
1255      Some(500),
1256      Some(2500),
1257      (Interval::new(Some(500), Some(2500)), Class::Body),
1258      false,
1259      true,
1260    );
1261
1262    // Restriction:       <---------------->
1263    // Request:   500------------------------------->
1264    // Result:    500------------------------------->
1265    test_interval_suppressed(
1266      None,
1267      None,
1268      Some(500),
1269      None,
1270      (Interval::new(Some(500), None), Class::Body),
1271      false,
1272      true,
1273    );
1274
1275    // Restriction:       <---------------->
1276    // Request:   <------------------------------2500
1277    // Result:    <------------------------------2500
1278    test_interval_suppressed(
1279      None,
1280      None,
1281      None,
1282      Some(2500),
1283      (Interval::new(None, Some(2500)), Class::Body),
1284      false,
1285      true,
1286    );
1287  }
1288
1289  #[test]
1290  fn validate_restrictions_format_none_allows_any() {
1291    let reference_restriction = ReferenceNameRestrictionBuilder::default()
1292      .name("chr1")
1293      .build()
1294      .unwrap();
1295    let rule = AuthorizationRuleBuilder::default()
1296      .location(test_location())
1297      .reference_name(reference_restriction)
1298      .build()
1299      .unwrap();
1300    let restrictions = AuthorizationRestrictionsBuilder::default()
1301      .rule(rule)
1302      .build()
1303      .unwrap();
1304
1305    let mut query = HashMap::new();
1306    query.insert("referenceName".to_string(), "chr1".to_string());
1307    query.insert("format".to_string(), "CRAM".to_string());
1308
1309    let request = create_test_query(Endpoint::Reads, "sample1", query);
1310    let result =
1311      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], false);
1312    assert!(result.is_ok());
1313  }
1314
1315  #[test]
1316  fn validate_restrictions_path_with_leading_slash() {
1317    let rule = AuthorizationRuleBuilder::default()
1318      .location(test_location())
1319      .build()
1320      .unwrap();
1321    let restrictions = AuthorizationRestrictionsBuilder::default()
1322      .rule(rule)
1323      .build()
1324      .unwrap();
1325    let request = create_test_query(Endpoint::Reads, "sample1", HashMap::new());
1326    let result =
1327      Auth::validate_restrictions(restrictions, request.id(), &mut [request.clone()], false);
1328    assert!(result.is_ok());
1329  }
1330
1331  #[tokio::test]
1332  async fn validate_authorization_missing_auth_header() {
1333    let mut auth = create_mock_auth_with_restrictions();
1334    let request = Request::new("sample1".to_string(), HashMap::new(), HeaderMap::new());
1335
1336    let result = auth.validate_jwt(request.headers()).await;
1337    assert!(result.is_err());
1338    assert!(matches!(
1339      result.unwrap_err(),
1340      HtsGetError::InvalidAuthentication(_)
1341    ));
1342  }
1343
1344  #[tokio::test]
1345  async fn validate_authorization_invalid_jwt_format() {
1346    let mut auth = create_mock_auth_with_restrictions();
1347    let request = create_request_with_auth_header("sample1", HashMap::new(), "invalid.jwt.token");
1348
1349    let result = auth.validate_jwt(request.headers()).await;
1350    assert!(result.is_err());
1351    assert!(matches!(
1352      result.unwrap_err(),
1353      HtsGetError::PermissionDenied(_)
1354    ));
1355  }
1356
1357  fn callout_with_forward(forward: Forward) -> Callout {
1358    Callout::new(
1359      Uri::from_static("https://www.example.com"),
1360      HttpClient::from(HttpClientConfig::default()),
1361      forward,
1362    )
1363  }
1364
1365  fn create_test_auth_config(public_key: Vec<u8>) -> AuthConfig {
1366    AuthConfigBuilder::default()
1367      .jwt_raw(JwtKey::PublicKey(public_key))
1368      .authorization(AuthorizationSourceBuilder::Callout(Box::new(Callout::new(
1369        Uri::from_static("https://www.example.com"),
1370        HttpClient::from(HttpClientConfig::default()),
1371        Forward::default(),
1372      ))))
1373      .build()
1374      .unwrap()
1375  }
1376
1377  fn create_test_query(endpoint: Endpoint, path: &str, query: HashMap<String, String>) -> Query {
1378    let request = Request::new(path.to_string(), query, HeaderMap::new());
1379    let format = match_format_from_query(&endpoint, request.query()).unwrap();
1380
1381    convert_to_query(request, format).unwrap()
1382  }
1383
1384  fn create_request_with_auth_header(
1385    path: &str,
1386    query: HashMap<String, String>,
1387    token: &str,
1388  ) -> Request {
1389    let mut headers = HeaderMap::new();
1390    headers.insert("authorization", format!("Bearer {token}").parse().unwrap());
1391    Request::new(path.to_string(), query, headers)
1392  }
1393
1394  fn create_mock_auth_with_restrictions() -> Auth {
1395    let (_, public_key) = generate_key_pair();
1396
1397    let config = create_test_auth_config(public_key);
1398    AuthBuilder::default().with_config(config).build().unwrap()
1399  }
1400
1401  fn test_interval_suppressed(
1402    restrict_start: Option<u32>,
1403    restrict_end: Option<u32>,
1404    request_start: Option<u32>,
1405    request_end: Option<u32>,
1406    expected_response: (Interval, Class),
1407    is_err: bool,
1408    suppress_interval: bool,
1409  ) {
1410    let mut reference_restriction = ReferenceNameRestrictionBuilder::default()
1411      .name("chr1")
1412      .format(Format::Bam);
1413
1414    if let Some(start) = restrict_start {
1415      reference_restriction = reference_restriction.start(start);
1416    }
1417    if let Some(end) = restrict_end {
1418      reference_restriction = reference_restriction.end(end);
1419    }
1420
1421    let reference_restriction = reference_restriction.build().unwrap();
1422    let rule = AuthorizationRuleBuilder::default()
1423      .location(test_location())
1424      .reference_name(reference_restriction)
1425      .build()
1426      .unwrap();
1427    let restrictions = AuthorizationRestrictionsBuilder::default()
1428      .rule(rule.clone())
1429      .build()
1430      .unwrap();
1431
1432    let mut query = HashMap::new();
1433    query.insert("referenceName".to_string(), "chr1".to_string());
1434    request_start.map(|start| query.insert("start".to_string(), start.to_string()));
1435    request_end.map(|end| query.insert("end".to_string(), end.to_string()));
1436
1437    let request = create_test_query(Endpoint::Reads, "sample1", query);
1438    let id = request.id().to_string();
1439    let mut slice = [request];
1440    let result = Auth::validate_restrictions(restrictions, &id, &mut slice, suppress_interval);
1441    if is_err {
1442      assert!(result.is_err());
1443    } else {
1444      assert!(result.is_ok());
1445    }
1446    assert_eq!(slice.first().unwrap().interval(), expected_response.0);
1447    assert_eq!(slice.last().unwrap().class(), expected_response.1);
1448  }
1449
1450  fn test_location() -> Location {
1451    Location::Simple(Box::new(SimpleLocation::new(
1452      Default::default(),
1453      "".to_string(),
1454      Some(PrefixOrId::Id("sample1".to_string())),
1455    )))
1456  }
1457}