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 auth_layer_with_rate_limit(State((auth, super::RateLimiter::disabled())), req, next).await
100}
101
102pub(crate) async fn auth_layer_with_rate_limit(
103 State((auth, failed_auth)): State<(AuthConfig, super::RateLimiter)>,
104 req: Request<Body>,
105 next: Next,
106) -> Response {
107 let path = req.uri().path();
108 if !auth.is_required()
109 || path == "/health"
110 || path == "/metrics"
111 || path == "/playground"
112 || req.method() == axum::http::Method::OPTIONS
113 {
114 return next.run(req).await;
115 }
116
117 let supplied = req
118 .headers()
119 .get(&AUTH_HEADER)
120 .and_then(|v| v.to_str().ok())
121 .and_then(|s| {
122 let (scheme, token) = s.split_once(' ')?;
123 scheme.eq_ignore_ascii_case("Bearer").then_some(token)
124 });
125
126 let ok = supplied.is_some_and(|token| auth.token_matches(token));
127
128 if !ok {
129 metrics::counter!("openkind_auth_failures_total", "transport" => "http").increment(1);
130 if let Some(peer) = req
131 .extensions()
132 .get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
133 {
134 if let Err(retry_after_ms) = failed_auth.check(peer.0.ip()) {
135 return crate::ApiError::RateLimited { retry_after_ms }.into_response();
136 }
137 }
138 let body = Json(serde_json::json!({
139 "error": {
140 "code": "unauthorized",
141 "message": "missing or invalid API key",
142 }
143 }));
144 let mut resp = (StatusCode::UNAUTHORIZED, body).into_response();
145 if let Ok(v) = HeaderValue::from_str("Bearer") {
146 resp.headers_mut()
147 .insert(axum::http::header::WWW_AUTHENTICATE, v);
148 }
149 if let Some(req_id) = req.extensions().get::<RequestId>() {
151 if let Ok(v) = HeaderValue::from_str(&req_id.0) {
152 resp.headers_mut().insert(REQUEST_ID_HEADER.clone(), v);
153 }
154 }
155 return resp;
156 }
157
158 next.run(req).await
159}
160
161fn digest_of(token: &str) -> [u8; 32] {
163 let digest = ring::digest::digest(&ring::digest::SHA256, token.as_bytes());
164 let mut out = [0u8; 32];
165 out.copy_from_slice(digest.as_ref());
166 out
167}
168
169pub fn secure_token_eq(a: &str, b: &str) -> bool {
175 use subtle::ConstantTimeEq;
176 digest_of(a).ct_eq(&digest_of(b)).into()
177}
178
179async fn auth_layer_dummy_handler() -> StatusCode {
181 StatusCode::OK
182}
183
184pub fn auth_layer_for(auth: AuthConfig) -> axum::Router {
186 axum::Router::new()
187 .route("/", axum::routing::get(auth_layer_dummy_handler))
189 .layer(axum::middleware::from_fn_with_state(auth, auth_layer))
190}