1use 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#[derive(Default, Debug)]
31pub struct AuthBuilder {
32 config: Option<AuthConfig>,
33}
34
35impl AuthBuilder {
36 pub fn with_config(mut self, config: AuthConfig) -> Self {
38 self.config = Some(config);
39 self
40 }
41
42 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#[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 pub fn config(&self) -> &AuthConfig {
82 &self.config
83 }
84
85 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 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 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 pub async fn decode_jwks(callout: &mut Callout, token: &str) -> HtsGetResult<DecodingKey> {
139 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 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 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 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 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 pub fn validate_restrictions(
246 restrictions: AuthorizationRestrictions,
247 path: &str,
248 queries: &mut [Query],
249 suppressed_interval: bool,
250 ) -> HtsGetResult<AuthorizationRestrictions> {
251 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 path.starts_with(&prefix)
262 }
263 PrefixOrId::Id(id) => {
264 id == path
266 }
267 }
268 }
269 Location::Regex(location) => {
270 location.regex().is_match(path)
272 }
273 _ => false,
275 }
276 })
277 .collect::<Vec<_>>();
278
279 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 for query in queries {
292 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 let name_match = restriction.reference_name().is_none()
303 || restriction.reference_name() == query.reference_name();
304 let format_match =
306 restriction.format().is_none() || restriction.format() == Some(query.format());
307 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); if suppressed_interval {
326 if allows_all.is_empty() && matching_restriction.is_none() {
327 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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}