1pub mod authz;
15#[cfg(feature = "builtin_jwt")]
16pub mod crypto;
17pub mod forward;
18#[cfg(feature = "builtin_jwt")]
19pub mod jwks;
20pub mod policy;
21#[cfg(feature = "builtin_jwt")]
22mod verifier;
23
24#[cfg(test)]
25mod tests;
26
27use std::collections::HashMap;
28use std::collections::HashSet;
29use std::sync::Arc;
30
31use axum::extract::State;
32use axum::http::header::{HeaderName, HeaderValue};
33use axum::http::{HeaderMap, StatusCode};
34use axum::middleware::Next;
35use axum::response::{IntoResponse, Response};
36use axum::Json;
37use serde_json::Value;
38
39use crate::config::{default_roles_claim, AuthConfig};
40use crate::hooks::TokenVerifier;
41use policy::Policies;
42
43enum Verifier {
45 #[cfg(feature = "builtin_jwt")]
50 Builtin(Box<verifier::ConfigVerifier>),
51 Injected(Arc<dyn TokenVerifier>),
53}
54
55pub struct Auth {
58 verifier: Verifier,
59 claims_headers: HashMap<String, String>,
60 roles_claim: String,
61 policies: Policies,
62}
63
64impl Auth {
65 pub fn build(
78 config: &AuthConfig,
79 verifier: Option<Arc<dyn TokenVerifier>>,
80 ) -> Result<Option<Arc<Self>>, String> {
81 if config.mode != "jwt" {
82 return Ok(None);
83 }
84
85 let verifier = match verifier {
86 Some(v) => {
87 if let Some(jwt) = &config.jwt {
88 if jwt.jwks_uri.is_some() || jwt.public_key_pem_file.is_some() {
89 tracing::warn!(
90 "auth.jwt names a key source, but an injected TokenVerifier is in use; \
91 the configured keys are ignored"
92 );
93 }
94 }
95 Verifier::Injected(v)
96 }
97 None => builtin_verifier(config)?,
98 };
99
100 let policies = match &config.forward_auth {
101 Some(fa) => Policies::compile(&fa.policies)?,
102 None => Policies::default(),
103 };
104
105 let (claims_headers, roles_claim) = match &config.jwt {
109 Some(jwt) => (jwt.claims_headers.clone(), jwt.roles_claim.clone()),
110 None => (HashMap::new(), default_roles_claim()),
111 };
112
113 Ok(Some(Arc::new(Self {
114 verifier,
115 claims_headers,
116 roles_claim,
117 policies,
118 })))
119 }
120
121 async fn verify(&self, token: &str) -> Option<Value> {
123 match &self.verifier {
124 #[cfg(feature = "builtin_jwt")]
125 Verifier::Builtin(v) => v.verify(token).await,
126 Verifier::Injected(v) => v.verify(token).await,
127 }
128 }
129}
130
131#[cfg(feature = "builtin_jwt")]
133fn builtin_verifier(config: &AuthConfig) -> Result<Verifier, String> {
134 let jwt = config
135 .jwt
136 .as_ref()
137 .ok_or("auth.mode is \"jwt\" but auth.jwt is not set")?;
138 Ok(Verifier::Builtin(Box::new(
139 verifier::ConfigVerifier::build(jwt)?,
140 )))
141}
142
143#[cfg(not(feature = "builtin_jwt"))]
146fn builtin_verifier(_config: &AuthConfig) -> Result<Verifier, String> {
147 Err(
148 "auth.mode is \"jwt\" but this build has no JWT crypto backend: enable the \
149 `rust_crypto` or `aws_lc_rs` feature, or inject a verifier with \
150 ProxyServer::with_token_verifier"
151 .to_string(),
152 )
153}
154
155pub(crate) enum AuthDecision {
157 Allow(HeaderMap, Option<Value>),
161 Unauthenticated(&'static str),
163 Forbidden(&'static str),
165}
166
167#[derive(Clone)]
172pub(crate) struct ValidatedClaims(pub(crate) std::sync::Arc<Value>);
173
174impl Auth {
175 pub(crate) async fn decide(
179 &self,
180 headers: &HeaderMap,
181 path: &str,
182 method: &str,
183 ) -> AuthDecision {
184 let claims = match bearer_token(headers) {
186 Some(token) => match self.verify(token).await {
187 Some(c) => Some(c),
188 None => return AuthDecision::Unauthenticated("invalid or expired token"),
189 },
190 None => None,
191 };
192
193 if let Some(policy) = self.policies.match_rule(path, method) {
194 if policy.require_auth && claims.is_none() {
195 return AuthDecision::Unauthenticated("authentication required");
196 }
197 if !policy.required_roles.is_empty() {
198 let Some(claims) = claims.as_ref() else {
201 return AuthDecision::Unauthenticated("authentication required");
202 };
203 let roles = extract_roles(claims, &self.roles_claim);
204 if !policy.required_roles.iter().all(|r| roles.contains(r)) {
205 return AuthDecision::Forbidden("insufficient role");
206 }
207 }
208 }
209
210 let mut claim_headers = HeaderMap::new();
211 if let Some(claims) = &claims {
212 inject_claim_headers(&mut claim_headers, claims, &self.claims_headers);
213 }
214 AuthDecision::Allow(claim_headers, claims)
215 }
216}
217
218pub async fn middleware(
220 State(auth): State<Arc<Auth>>,
221 mut request: axum::extract::Request,
222 next: Next,
223) -> Response {
224 let path = request.uri().path().to_string();
225 let method = request.method().as_str().to_ascii_uppercase();
226
227 strip_claim_headers(request.headers_mut(), &auth.claims_headers);
231
232 match auth.decide(request.headers(), &path, &method).await {
233 AuthDecision::Unauthenticated(msg) => unauthorized(msg),
234 AuthDecision::Forbidden(msg) => forbidden(msg),
235 AuthDecision::Allow(claim_headers, claims) => {
236 let dst = request.headers_mut();
237 for (name, value) in &claim_headers {
238 dst.insert(name.clone(), value.clone());
239 }
240 if let Some(claims) = claims {
243 request
244 .extensions_mut()
245 .insert(ValidatedClaims(std::sync::Arc::new(claims)));
246 }
247 next.run(request).await
248 }
249 }
250}
251
252fn bearer_token(headers: &HeaderMap) -> Option<&str> {
257 let value = headers.get("authorization")?.to_str().ok()?;
258 let token = value
259 .strip_prefix("Bearer ")
260 .or_else(|| value.strip_prefix("bearer "))?;
261 let token = token.trim();
262 (!token.is_empty()).then_some(token)
263}
264
265fn claim_at<'a>(claims: &'a Value, path: &str) -> Option<&'a Value> {
267 let mut cur = claims;
268 for seg in path.split('.') {
269 cur = cur.get(seg)?;
270 }
271 Some(cur)
272}
273
274fn extract_roles(claims: &Value, roles_claim: &str) -> HashSet<String> {
276 claim_at(claims, roles_claim)
277 .and_then(Value::as_array)
278 .map(|arr| {
279 arr.iter()
280 .filter_map(|v| v.as_str().map(str::to_string))
281 .collect()
282 })
283 .unwrap_or_default()
284}
285
286fn strip_claim_headers(headers: &mut HeaderMap, mapping: &HashMap<String, String>) {
289 for header in mapping.values() {
290 if let Ok(name) = HeaderName::try_from(header.as_str()) {
291 while headers.remove(&name).is_some() {}
292 }
293 }
294}
295
296fn inject_claim_headers(
298 headers: &mut HeaderMap,
299 claims: &Value,
300 mapping: &HashMap<String, String>,
301) {
302 for (claim, header) in mapping {
303 let Some(value) = claim_at(claims, claim) else {
304 continue;
305 };
306 let rendered = match value {
307 Value::String(s) => s.clone(),
308 Value::Number(n) => n.to_string(),
309 Value::Bool(b) => b.to_string(),
310 _ => continue,
312 };
313 if let (Ok(name), Ok(val)) = (
314 HeaderName::try_from(header.as_str()),
315 HeaderValue::try_from(rendered),
316 ) {
317 headers.insert(name, val);
318 }
319 }
320}
321
322fn unauthorized(message: &str) -> Response {
323 (
324 StatusCode::UNAUTHORIZED,
325 Json(serde_json::json!({ "error": "UNAUTHENTICATED", "message": message })),
326 )
327 .into_response()
328}
329
330fn forbidden(message: &str) -> Response {
331 (
332 StatusCode::FORBIDDEN,
333 Json(serde_json::json!({ "error": "PERMISSION_DENIED", "message": message })),
334 )
335 .into_response()
336}