Skip to main content

haste_server/auth_n/middleware/
jwt.rs

1use crate::{
2    auth_n::certificates,
3    config::ServerConfig,
4    extract::{
5        bearer_token::AuthBearer,
6        path_tenant::{ProjectIdentifier, TenantIdentifier},
7    },
8    route_path::{api_fhir_root_url, project_path},
9    services::ServerState,
10};
11use axum::{
12    extract::{OriginalUri, Request, State},
13    http::{HeaderMap, StatusCode, Uri},
14    middleware::Next,
15    response::{IntoResponse as _, Response},
16};
17use axum_extra::extract::Cached;
18use derivative::Derivative;
19use haste_fhir_model::r4::generated::terminology::IssueType;
20use haste_fhir_operation_error::OperationOutcomeError;
21use haste_fhir_search::SearchEngine;
22use haste_fhir_terminology::FHIRTerminology;
23use haste_jwt::{ProjectId, TenantId, claims::UserTokenClaims};
24use haste_repository::Repository;
25use jsonwebtoken::Validation;
26use std::{path::PathBuf, sync::Arc};
27use url::Url;
28
29#[derive(Derivative)]
30#[derivative(Debug)]
31pub struct User {
32    #[allow(dead_code)]
33    #[derivative(Debug = "ignore")]
34    pub token: Option<String>,
35    pub claims: haste_jwt::claims::UserTokenClaims,
36}
37
38fn validate_jwt(
39    config: &ServerConfig,
40    tenant: &TenantId,
41    project: &ProjectId,
42    token: &str,
43) -> Result<UserTokenClaims, StatusCode> {
44    let header = jsonwebtoken::decode_header(token).map_err(|_| StatusCode::UNAUTHORIZED)?;
45
46    let kid = header.kid.ok_or(StatusCode::UNAUTHORIZED)?;
47
48    let cert_provider = certificates::get_certification_provider(config);
49
50    let decoding_key = cert_provider
51        .decoding_key(kid.as_str())
52        .map_err(|_| StatusCode::UNAUTHORIZED)?;
53
54    // Per SMART on FHIR, `aud` identifies the FHIR resource server a token was
55    // minted for, binding the token to this tenant/project's FHIR endpoint.
56    let expected_audience = api_fhir_root_url(&config.api_uri, tenant, project)
57        .map_err(|_| StatusCode::UNAUTHORIZED)?;
58
59    let mut validation = Validation::new(jsonwebtoken::Algorithm::RS256);
60    validation.set_audience(&[expected_audience.as_str()]);
61
62    let result =
63        jsonwebtoken::decode::<UserTokenClaims>(token, &decoding_key.decoding_key, &validation)
64            .map_err(|_| StatusCode::UNAUTHORIZED)?;
65
66    Ok(result.claims)
67}
68
69pub fn derive_well_known_openid_configuration_url(
70    api_url: &str,
71    tenant: &TenantId,
72    project: &ProjectId,
73) -> Result<Url, OperationOutcomeError> {
74    let path = PathBuf::from("/.well-known/openid-configuration");
75
76    if let Ok(api_url) = Url::parse(api_url) {
77        api_url
78            .join(
79                path.join(project_path(tenant, project).strip_prefix("/").unwrap())
80                    .to_str()
81                    .unwrap_or_default(),
82            )
83            .map_err(|e| {
84                tracing::error!("Failed to derive well-known URL: {:?}", e);
85                OperationOutcomeError::error(
86                    IssueType::invalid(),
87                    "Invalid API URL configured".to_string(),
88                )
89            })
90    } else {
91        Err(OperationOutcomeError::error(
92            IssueType::invalid(),
93            "Invalid API URL configured".to_string(),
94        ))
95    }
96}
97
98pub fn derive_protected_resource_metadata_url(
99    resource_uri: &Uri,
100    api_url: &str,
101) -> Result<Url, OperationOutcomeError> {
102    let path = PathBuf::from("/.well-known/oauth-protected-resource");
103    if let Ok(api_url) = Url::parse(api_url) {
104        let tenant_url = api_url
105            .join(
106                path.join(resource_uri.path().strip_prefix("/").unwrap_or_default())
107                    .to_str()
108                    .unwrap_or_default(),
109            )
110            .map_err(|e| {
111                tracing::error!("Failed to derive well-known URL: {:?}", e);
112                OperationOutcomeError::error(
113                    IssueType::invalid(),
114                    "Invalid API URL configured".to_string(),
115                )
116            })?;
117
118        Ok(tenant_url)
119    } else {
120        Err(OperationOutcomeError::error(
121            IssueType::invalid(),
122            "Invalid API URL configured".to_string(),
123        ))
124    }
125}
126
127fn invalid_jwt_response(uri: &Uri, api_url: &str, status_code: StatusCode) -> Response {
128    tracing::warn!(
129        "Invalid JWT token provided in request sending '{}'",
130        status_code
131    );
132
133    let Ok(protected_resource_metadata_url) = derive_protected_resource_metadata_url(uri, api_url)
134    else {
135        return (status_code).into_response();
136    };
137
138    let mut headers = HeaderMap::new();
139    headers.insert(
140        axum::http::header::WWW_AUTHENTICATE,
141        format!(
142            r#"Bearer resource_metadata="{}""#,
143            protected_resource_metadata_url
144        )
145        .parse()
146        .unwrap(),
147    );
148    (status_code, headers).into_response()
149}
150
151pub async fn token_verifcation<
152    Repo: Repository + Send + Sync + 'static,
153    Search: SearchEngine + Send + Sync + 'static,
154    Terminology: FHIRTerminology + Send + Sync + 'static,
155>(
156    State(state): State<Arc<ServerState<Repo, Search, Terminology>>>,
157    Cached(TenantIdentifier { tenant }): Cached<TenantIdentifier>,
158    Cached(ProjectIdentifier { project }): Cached<ProjectIdentifier>,
159    // run the `HeaderMap` extractor
160    AuthBearer(token): AuthBearer,
161    // you can also add more extractors here but the last
162    // extractor must implement `FromRequest` which
163    // `Request` does
164    OriginalUri(uri): OriginalUri,
165    mut request: Request,
166    next: Next,
167) -> Result<Response, Response> {
168    let Some(token) = token else {
169        return Err(invalid_jwt_response(
170            &uri,
171            &state.config.api_uri,
172            StatusCode::UNAUTHORIZED,
173        ));
174    };
175
176    match validate_jwt(state.config.as_ref(), &tenant, &project, &token) {
177        Ok(claims) => {
178            request.extensions_mut().insert(Arc::new(User {
179                token: Some(token),
180                claims,
181            }));
182            Ok(next.run(request).await)
183        }
184        Err(status_code) => match status_code {
185            StatusCode::UNAUTHORIZED => Err(invalid_jwt_response(
186                &uri,
187                &state.config.api_uri,
188                status_code,
189            )),
190            _ => Err((status_code).into_response()),
191        },
192    }
193}