Skip to main content

haste_server/auth_n/oidc/middleware/
oidc_parameter_inject.rs

1use axum::RequestExt;
2use axum::http::{Method, StatusCode};
3use axum::response::IntoResponse;
4use axum::{body::Body, extract::Request, response::Response};
5use axum::{body::to_bytes, extract::Query};
6use haste_fhir_model::r4::generated::terminology::IssueType;
7use haste_fhir_operation_error::OperationOutcomeError;
8use serde::Deserialize;
9use std::sync::Arc;
10use std::task::{Context, Poll};
11use std::{collections::HashMap, pin::Pin};
12use tower::{Layer, Service};
13
14#[derive(Deserialize, Clone, Debug)]
15pub struct OIDCParameters {
16    pub parameters: HashMap<String, String>,
17    pub launch_parameters: Option<HashMap<String, String>>,
18}
19
20#[derive(Clone, Debug)]
21pub struct ParameterConfig {
22    pub required_parameters: Vec<String>,
23    pub optional_parameters: Vec<String>,
24    pub allow_launch_parameters: bool,
25}
26
27#[derive(Clone)]
28pub struct OIDCParameterInjectLayer {
29    state: Arc<ParameterConfig>,
30}
31
32impl<S> Layer<S> for OIDCParameterInjectLayer {
33    type Service = OIDCParameterInjectService<S>;
34
35    fn layer(&self, inner: S) -> Self::Service {
36        OIDCParameterInjectService {
37            inner,
38            state: self.state.clone(),
39        }
40    }
41}
42
43impl OIDCParameterInjectLayer {
44    pub fn new(state: Arc<ParameterConfig>) -> Self {
45        OIDCParameterInjectLayer { state }
46    }
47}
48
49#[derive(Clone)]
50pub struct OIDCParameterInjectService<S> {
51    inner: S,
52    state: Arc<ParameterConfig>,
53}
54
55fn validate_parameter(param_name: &str, param_value: &str) -> Result<(), String> {
56    if param_value.is_empty() {
57        return Err("Parameter cannot be empty".to_string());
58    }
59
60    if param_name == "response_type" && !["code", "token"].contains(&param_value) {
61        return Err(format!("Invalid response_type: {}", param_value));
62    }
63
64    Ok(())
65}
66
67impl<'a, T> Service<Request<Body>> for OIDCParameterInjectService<T>
68where
69    T: Service<Request, Response = Response> + Send + 'static + Clone,
70    T::Future: Send + 'static,
71    T::Error: IntoResponse,
72{
73    type Response = T::Response;
74    type Error = T::Error;
75    type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
76
77    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
78        self.inner.poll_ready(cx)
79    }
80
81    fn call(&mut self, mut request: Request) -> Self::Future {
82        // https://docs.rs/tower/latest/tower/trait.Service.html#be-careful-when-cloning-inner-services
83        let clone = self.inner.clone();
84        // take the service that was ready
85        let mut inner = std::mem::replace(&mut self.inner, clone);
86        let parameter_config = self.state.clone();
87
88        Box::pin(async move {
89            let Ok(Query(query_params)) = request
90                .extract_parts::<Query<HashMap<String, String>>>()
91                .await
92            else {
93                return Ok((StatusCode::BAD_REQUEST, "".to_string()).into_response());
94            };
95
96            let (parts, body) = request.into_parts();
97            let bytes = to_bytes(body, 10000).await;
98            let Ok(bytes) = bytes else {
99                return Ok((
100                    StatusCode::BAD_REQUEST,
101                    "Body was to large size limit 10k bytes".to_string(),
102                )
103                    .into_response());
104            };
105
106            let content_type = parts
107                .headers
108                .get(axum::http::header::CONTENT_TYPE)
109                .and_then(|v| v.to_str().ok())
110                .unwrap_or("");
111
112            // Either check the body if serializes or check the query params.
113            let unvalidated_parameters = match parts.method {
114                Method::POST => match content_type {
115                    "application/x-www-form-urlencoded" => {
116                        let mut form_params =
117                            serde_html_form::from_bytes::<HashMap<String, String>>(&bytes)
118                                .unwrap_or_else(|_e| HashMap::new());
119                        form_params.extend(query_params);
120
121                        Ok(form_params)
122                    }
123                    "application/json" => {
124                        let mut json_params =
125                            serde_json::from_slice::<HashMap<String, String>>(&bytes)
126                                .unwrap_or_else(|_e| HashMap::new());
127                        json_params.extend(query_params);
128                        Ok(json_params)
129                    }
130                    content_type => Err(OperationOutcomeError::error(
131                        IssueType::not_supported(),
132                        format!(
133                            "Unsupported Content-Type for OIDC parameter injection '{}'",
134                            content_type
135                        ),
136                    )),
137                },
138                Method::GET => Ok(query_params),
139                _ => Err(OperationOutcomeError::error(
140                    IssueType::not_supported(),
141                    "Unsupported HTTP method for OIDC parameter injection".to_string(),
142                )),
143            };
144
145            let Ok(unvalidated_parameters) = unvalidated_parameters else {
146                let error = unvalidated_parameters.err().unwrap();
147                return Ok(error.into_response());
148            };
149
150            let mut oidc_parameters = OIDCParameters {
151                parameters: HashMap::new(),
152                launch_parameters: None,
153            };
154
155            // Check for required parameters
156            for (is_required, parameter_name) in parameter_config
157                .required_parameters
158                .iter()
159                .map(|p| (true, p))
160                .chain(
161                    parameter_config
162                        .optional_parameters
163                        .iter()
164                        .map(|p| (false, p)),
165                )
166            {
167                if let Some(parameter_value) = unvalidated_parameters.get(parameter_name) {
168                    if let Err(_e) = validate_parameter(parameter_name, parameter_value) {
169                        return Ok((
170                            StatusCode::BAD_REQUEST,
171                            format!("Invalid parameter: '{}'", parameter_name),
172                        )
173                            .into_response());
174                    }
175
176                    oidc_parameters
177                        .parameters
178                        .insert(parameter_name.clone(), parameter_value.clone());
179                } else if is_required {
180                    return Ok((
181                        StatusCode::BAD_REQUEST,
182                        format!("Missing required parameter: '{parameter_name}'",),
183                    )
184                        .into_response());
185                }
186            }
187            // Launch parameters are for SMART apps e.g. launch/patient
188            if parameter_config.allow_launch_parameters {
189                let mut launch_parameters = HashMap::new();
190                for launch_param_name in unvalidated_parameters
191                    .keys()
192                    .filter(|k| k.starts_with("launch/"))
193                {
194                    let launch_param_value = unvalidated_parameters
195                        .get(launch_param_name)
196                        .cloned()
197                        .unwrap_or_default();
198
199                    let parts = launch_param_name.split('/').collect::<Vec<_>>();
200                    if parts.len() != 2 {
201                        return Ok((
202                            StatusCode::BAD_REQUEST,
203                            format!("Invalid launch parameter: '{launch_param_name}'"),
204                        )
205                            .into_response());
206                    }
207
208                    launch_parameters.insert(launch_param_name.to_string(), launch_param_value);
209                }
210                oidc_parameters.launch_parameters = Some(launch_parameters);
211            }
212
213            let new_body = Body::from(bytes);
214            let mut new_request = Request::from_parts(parts, new_body);
215            new_request.extensions_mut().insert(oidc_parameters);
216
217            let future = inner.call(new_request);
218            let response: Response = future.await?;
219            Ok(response)
220        })
221    }
222}