openkind_api/middleware/
auth.rs1use std::sync::Arc;
4
5use axum::{
6 body::Body,
7 extract::State,
8 http::{HeaderName, HeaderValue, Request, StatusCode},
9 middleware::Next,
10 response::{IntoResponse, Response},
11 Json,
12};
13
14use super::request_id::{RequestId, REQUEST_ID_HEADER};
15
16pub const AUTH_HEADER: HeaderName = HeaderName::from_static("authorization");
18
19#[derive(Clone, Default)]
21pub struct AuthConfig {
22 pub expected: Arc<Option<String>>,
24 expected_digest: Option<(Arc<str>, [u8; 32])>,
27}
28
29impl AuthConfig {
30 pub fn new(expected: Option<String>) -> Self {
32 let expected_digest = expected
33 .as_deref()
34 .map(|token| (Arc::from(token), digest_of(token)));
35 Self {
36 expected: Arc::new(expected),
37 expected_digest,
38 }
39 }
40
41 pub fn resolve_api_key_with<F>(get_env: F) -> Option<String>
46 where
47 F: Fn(&str) -> Result<String, std::env::VarError>,
48 {
49 get_env("OPENKIND_API_KEY")
50 .ok()
51 .filter(|s| !s.is_empty())
52 .or_else(|| {
53 get_env("OPENDECISION_API_KEY")
54 .ok()
55 .filter(|s| !s.is_empty())
56 })
57 .or_else(|| get_env("TYPESAFE_API_KEY").ok().filter(|s| !s.is_empty()))
58 .or_else(|| get_env("OPENPICK_API_KEY").ok().filter(|s| !s.is_empty()))
59 }
60
61 pub fn from_env() -> Self {
63 Self::new(Self::resolve_api_key_with(|k| std::env::var(k)))
64 }
65
66 pub fn is_required(&self) -> bool {
68 self.expected.is_some()
69 }
70
71 pub(crate) fn token_matches(&self, supplied: &str) -> bool {
72 use subtle::ConstantTimeEq;
73 let Some(expected) = self.expected.as_deref() else {
74 return false;
75 };
76 let expected_digest = self
79 .expected_digest
80 .as_ref()
81 .filter(|(cached, _)| cached.as_ref() == expected)
82 .map(|(_, digest)| *digest)
83 .unwrap_or_else(|| digest_of(expected));
84 digest_of(supplied).ct_eq(&expected_digest).into()
85 }
86}
87
88pub async fn auth_layer(
95 State(auth): State<AuthConfig>,
96 req: Request<Body>,
97 next: Next,
98) -> Response {
99 let path = req.uri().path();
100 if !auth.is_required()
101 || path == "/health"
102 || path == "/metrics"
103 || path == "/playground"
104 || req.method() == axum::http::Method::OPTIONS
105 {
106 return next.run(req).await;
107 }
108
109 let supplied = req
110 .headers()
111 .get(&AUTH_HEADER)
112 .and_then(|v| v.to_str().ok())
113 .and_then(|s| {
114 let (scheme, token) = s.split_once(' ')?;
115 scheme.eq_ignore_ascii_case("Bearer").then_some(token)
116 });
117
118 let ok = supplied.is_some_and(|token| auth.token_matches(token));
119
120 if !ok {
121 let body = Json(serde_json::json!({
122 "error": {
123 "code": "unauthorized",
124 "message": "missing or invalid API key",
125 }
126 }));
127 let mut resp = (StatusCode::UNAUTHORIZED, body).into_response();
128 if let Ok(v) = HeaderValue::from_str("Bearer") {
129 resp.headers_mut()
130 .insert(axum::http::header::WWW_AUTHENTICATE, v);
131 }
132 if let Some(req_id) = req.extensions().get::<RequestId>() {
134 if let Ok(v) = HeaderValue::from_str(&req_id.0) {
135 resp.headers_mut().insert(REQUEST_ID_HEADER.clone(), v);
136 }
137 }
138 return resp;
139 }
140
141 next.run(req).await
142}
143
144fn digest_of(token: &str) -> [u8; 32] {
146 let digest = ring::digest::digest(&ring::digest::SHA256, token.as_bytes());
147 let mut out = [0u8; 32];
148 out.copy_from_slice(digest.as_ref());
149 out
150}
151
152pub fn secure_token_eq(a: &str, b: &str) -> bool {
158 use subtle::ConstantTimeEq;
159 digest_of(a).ct_eq(&digest_of(b)).into()
160}
161
162async fn auth_layer_dummy_handler() -> StatusCode {
164 StatusCode::OK
165}
166
167pub fn auth_layer_for(auth: AuthConfig) -> axum::Router {
169 axum::Router::new()
170 .route("/", axum::routing::get(auth_layer_dummy_handler))
172 .layer(axum::middleware::from_fn_with_state(auth, auth_layer))
173}