haste_server/auth_n/oidc/middleware/
oidc_parameter_inject.rs1use 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(¶m_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 let clone = self.inner.clone();
84 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 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 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 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}