1use axum::{
2 extract::Request,
3 http::{HeaderValue, Method, StatusCode},
4 middleware::Next,
5 response::{IntoResponse, Response},
6 Json,
7};
8use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
9use hmac::{Hmac, Mac};
10use serde::Serialize;
11use sha2::Sha256;
12use std::collections::HashSet;
13use std::sync::Arc;
14use std::time::{SystemTime, UNIX_EPOCH};
15use subtle::ConstantTimeEq;
16use tracing::warn;
17
18const GALLERY_MEDIA_TOKEN_CONTEXT: &[u8] = b"mold-gallery-media-v2\nGET\n";
19pub(crate) const GALLERY_MEDIA_TOKEN_TTL_SECS: u64 = 15 * 60;
20const GALLERY_SIGNING_SECRET_BYTES: usize = 32;
21
22type HmacSha256 = Hmac<Sha256>;
23
24pub type AuthState = Option<Arc<ApiKeySet>>;
26
27pub struct ApiKeySet {
29 keys: HashSet<String>,
30 gallery_signing_secret: [u8; GALLERY_SIGNING_SECRET_BYTES],
31}
32
33impl ApiKeySet {
34 pub fn new(keys: HashSet<String>) -> Self {
35 Self::try_new(keys).expect("OS randomness is required for gallery media authentication")
36 }
37
38 fn try_new(keys: HashSet<String>) -> Result<Self, getrandom::Error> {
39 let mut gallery_signing_secret = [0_u8; GALLERY_SIGNING_SECRET_BYTES];
40 getrandom::fill(&mut gallery_signing_secret)?;
41 Ok(Self {
42 keys,
43 gallery_signing_secret,
44 })
45 }
46
47 #[cfg(test)]
48 pub(crate) fn new_with_gallery_signing_secret(
49 keys: HashSet<String>,
50 gallery_signing_secret: [u8; GALLERY_SIGNING_SECRET_BYTES],
51 ) -> Self {
52 Self {
53 keys,
54 gallery_signing_secret,
55 }
56 }
57
58 pub fn contains(&self, candidate: &str) -> bool {
59 let candidate_bytes = candidate.as_bytes();
62 let mut found = subtle::Choice::from(0u8);
63 for k in &self.keys {
64 found |= k.as_bytes().ct_eq(candidate_bytes);
65 }
66 found.into()
67 }
68
69 fn validates_gallery_media_token(
70 &self,
71 token: &str,
72 media_path: &str,
73 expires_at: u64,
74 now: u64,
75 ) -> bool {
76 if now >= expires_at {
77 return false;
78 }
79
80 let Ok(signature) = URL_SAFE_NO_PAD.decode(token) else {
81 return false;
82 };
83 if signature.len() != 32 {
84 return false;
85 }
86
87 let mac = gallery_media_mac(&self.gallery_signing_secret, media_path, expires_at);
88 mac.verify_slice(&signature).is_ok()
89 }
90
91 pub(crate) fn issue_gallery_media_token(&self, media_path: &str) -> (String, u64) {
92 let expires_at = unix_timestamp().saturating_add(GALLERY_MEDIA_TOKEN_TTL_SECS);
93 (
94 sign_gallery_media_token(&self.gallery_signing_secret, media_path, expires_at),
95 expires_at,
96 )
97 }
98
99 #[cfg(test)]
100 pub(crate) fn sign_gallery_media_token_for_tests(
101 &self,
102 media_path: &str,
103 expires_at: u64,
104 ) -> String {
105 sign_gallery_media_token(&self.gallery_signing_secret, media_path, expires_at)
106 }
107}
108
109#[derive(Clone, Copy)]
112pub(crate) struct ApiKeyAuthenticated;
113
114#[derive(Debug, Serialize)]
115struct AuthError {
116 error: String,
117 code: String,
118}
119
120pub fn load_api_keys() -> anyhow::Result<AuthState> {
129 let raw = match std::env::var("MOLD_API_KEY") {
130 Ok(v) if !v.is_empty() => v,
131 _ => return Ok(None),
132 };
133
134 let keys = if let Some(path) = raw.strip_prefix('@') {
135 let contents = std::fs::read_to_string(path)
136 .map_err(|e| anyhow::anyhow!("failed to read API key file {path}: {e}"))?;
137 contents
138 .lines()
139 .map(str::trim)
140 .filter(|l| !l.is_empty() && !l.starts_with('#'))
141 .map(String::from)
142 .collect::<HashSet<_>>()
143 } else {
144 raw.split(',')
145 .map(str::trim)
146 .filter(|s| !s.is_empty())
147 .map(String::from)
148 .collect::<HashSet<_>>()
149 };
150
151 if keys.is_empty() {
152 return Ok(None);
153 }
154
155 tracing::info!(num_keys = keys.len(), "API key authentication enabled");
156 let key_set = ApiKeySet::try_new(keys)
157 .map_err(|error| anyhow::anyhow!("failed to generate gallery signing secret: {error}"))?;
158 Ok(Some(Arc::new(key_set)))
159}
160
161const EXEMPT_PATHS: &[&str] = &["/health", "/api/docs", "/api/openapi.json"];
163
164pub async fn require_api_key(request: Request, next: Next) -> Response {
166 let mut request = request;
167
168 let auth_state = request.extensions().get::<AuthState>().cloned();
170
171 let key_set = match auth_state.as_ref().and_then(|s| s.as_ref()) {
172 Some(ks) => ks,
173 None => return next.run(request).await, };
175
176 let path = request.uri().path();
178 if EXEMPT_PATHS.contains(&path) {
179 return next.run(request).await;
180 }
181
182 if is_gallery_image_read(&request) {
187 if let Some((token, expires_at)) = gallery_media_ticket_query(&request) {
188 let now = unix_timestamp();
189 if key_set.validates_gallery_media_token(token, path, expires_at, now) {
190 return next.run(request).await;
191 }
192 if request.headers().get("x-api-key").is_none() {
193 warn!(path = %path, "rejected request with invalid or expired gallery media token");
194 return unauthorized("invalid or expired gallery media token");
195 }
196 }
197 }
198
199 match request.headers().get("x-api-key") {
201 Some(value) => {
202 let candidate = value.to_str().unwrap_or("");
203 if key_set.contains(candidate) {
204 request.extensions_mut().insert(ApiKeyAuthenticated);
205 next.run(request).await
206 } else {
207 warn!("rejected request with invalid API key");
208 unauthorized("invalid API key")
209 }
210 }
211 None => {
212 warn!(path = %path, "rejected request without API key");
213 unauthorized("missing X-Api-Key header")
214 }
215 }
216}
217
218fn is_gallery_image_read(request: &Request) -> bool {
219 if request.method() != Method::GET && request.method() != Method::HEAD {
220 return false;
221 }
222
223 is_gallery_image_path(request.uri().path())
224}
225
226pub(crate) fn is_gallery_image_path(path: &str) -> bool {
227 if path.contains('?') || path.contains('#') {
228 return false;
229 }
230
231 path.strip_prefix("/api/gallery/image/")
232 .is_some_and(|filename| !filename.is_empty() && !filename.contains('/'))
233}
234
235fn gallery_media_ticket_query(request: &Request) -> Option<(&str, u64)> {
236 let mut token = None;
237 let mut expires_at = None;
238
239 for pair in request.uri().query()?.split('&') {
240 let (name, value) = pair.split_once('=').unwrap_or((pair, ""));
241 match name {
242 "media_token" if token.is_none() => token = Some(value),
243 "expires" if expires_at.is_none() => expires_at = value.parse::<u64>().ok(),
244 "media_token" | "expires" => return None,
247 _ => {}
248 }
249 }
250
251 Some((token?, expires_at?))
252}
253
254fn gallery_media_mac(signing_secret: &[u8], media_path: &str, expires_at: u64) -> HmacSha256 {
255 let mut mac =
256 HmacSha256::new_from_slice(signing_secret).expect("HMAC-SHA256 accepts keys of any length");
257 mac.update(GALLERY_MEDIA_TOKEN_CONTEXT);
258 mac.update(media_path.as_bytes());
259 mac.update(b"\n");
260 mac.update(expires_at.to_string().as_bytes());
261 mac
262}
263
264fn sign_gallery_media_token(signing_secret: &[u8], media_path: &str, expires_at: u64) -> String {
265 let signature = gallery_media_mac(signing_secret, media_path, expires_at)
266 .finalize()
267 .into_bytes();
268 URL_SAFE_NO_PAD.encode(signature)
269}
270
271fn unix_timestamp() -> u64 {
272 SystemTime::now()
273 .duration_since(UNIX_EPOCH)
274 .unwrap_or_default()
275 .as_secs()
276}
277
278fn unauthorized(msg: &str) -> Response {
279 let body = AuthError {
280 error: msg.to_string(),
281 code: "UNAUTHORIZED".to_string(),
282 };
283 (StatusCode::UNAUTHORIZED, Json(body)).into_response()
284}
285
286pub async fn inject_auth_state(
288 axum::extract::State(auth): axum::extract::State<AuthState>,
289 mut request: Request,
290 next: Next,
291) -> Response {
292 request.extensions_mut().insert(auth);
293 next.run(request).await
294}
295
296pub fn api_key_header_name() -> HeaderValue {
300 HeaderValue::from_static("x-api-key")
301}
302
303#[cfg(test)]
304mod tests {
305 use super::*;
306
307 fn env_lock() -> &'static std::sync::Mutex<()> {
309 static LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
310 &LOCK
311 }
312
313 #[test]
314 fn parse_single_key() {
315 let _lock = env_lock().lock().unwrap();
316 unsafe { std::env::set_var("MOLD_API_KEY", "secret123") };
317 let state = load_api_keys().unwrap();
318 unsafe { std::env::remove_var("MOLD_API_KEY") };
319 let ks = state.as_ref().unwrap();
320 assert!(ks.contains("secret123"));
321 assert!(!ks.contains("wrong"));
322 }
323
324 #[test]
325 fn parse_comma_separated_keys() {
326 let _lock = env_lock().lock().unwrap();
327 unsafe { std::env::set_var("MOLD_API_KEY", "key1,key2, key3 ") };
328 let state = load_api_keys().unwrap();
329 unsafe { std::env::remove_var("MOLD_API_KEY") };
330 let ks = state.as_ref().unwrap();
331 assert!(ks.contains("key1"));
332 assert!(ks.contains("key2"));
333 assert!(ks.contains("key3"));
334 assert!(!ks.contains("key4"));
335 }
336
337 #[test]
338 fn parse_file_keys() {
339 let _lock = env_lock().lock().unwrap();
340 let dir = std::env::temp_dir().join("mold_test_keys");
341 let _ = std::fs::create_dir_all(&dir);
342 let path = dir.join("keys.txt");
343 std::fs::write(&path, "alpha\n# comment\n\nbeta\n").unwrap();
344 let env_val = format!("@{}", path.display());
345 unsafe { std::env::set_var("MOLD_API_KEY", &env_val) };
346 let state = load_api_keys().unwrap();
347 unsafe { std::env::remove_var("MOLD_API_KEY") };
348 let _ = std::fs::remove_file(&path);
349 let ks = state.as_ref().unwrap();
350 assert!(ks.contains("alpha"));
351 assert!(ks.contains("beta"));
352 assert!(!ks.contains("# comment"));
353 }
354
355 #[test]
356 fn empty_env_returns_none() {
357 let _lock = env_lock().lock().unwrap();
358 unsafe { std::env::set_var("MOLD_API_KEY", "") };
359 let state = load_api_keys().unwrap();
360 unsafe { std::env::remove_var("MOLD_API_KEY") };
361 assert!(state.is_none());
362 }
363
364 #[test]
365 fn unset_env_returns_none() {
366 let _lock = env_lock().lock().unwrap();
367 unsafe { std::env::remove_var("MOLD_API_KEY") };
368 let state = load_api_keys().unwrap();
369 assert!(state.is_none());
370 }
371
372 #[test]
373 fn constant_time_comparison_rejects_wrong_key() {
374 let ks = ApiKeySet::new(HashSet::from(["correct-key".to_string()]));
375 assert!(ks.contains("correct-key"));
376 assert!(!ks.contains("wrong-key"));
377 assert!(!ks.contains(""));
378 }
379
380 #[test]
381 fn gallery_media_token_is_url_safe_and_valid_until_expiry() {
382 const MEDIA_PATH: &str = "/api/gallery/image/clip.mp4";
383 let ks = ApiKeySet::new_with_gallery_signing_secret(
384 HashSet::from(["correct-key".to_string()]),
385 [0x42; GALLERY_SIGNING_SECRET_BYTES],
386 );
387 let token = ks.sign_gallery_media_token_for_tests(MEDIA_PATH, 1_900);
388
389 assert_eq!(token.len(), 43);
390 assert!(token
391 .bytes()
392 .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-' || byte == b'_'));
393 assert!(ks.validates_gallery_media_token(&token, MEDIA_PATH, 1_900, 1_000));
394 assert!(!ks.validates_gallery_media_token(&token, MEDIA_PATH, 1_900, 1_900));
395 }
396
397 #[test]
398 fn gallery_media_token_rejects_tampering_and_other_signing_secrets() {
399 const MEDIA_PATH: &str = "/api/gallery/image/clip.mp4";
400 let ks = ApiKeySet::new_with_gallery_signing_secret(
401 HashSet::from(["correct-key".to_string()]),
402 [0x42; GALLERY_SIGNING_SECRET_BYTES],
403 );
404 let token = ks.sign_gallery_media_token_for_tests(MEDIA_PATH, 1_900);
405 let mut tampered = token.clone();
406 tampered.replace_range(..1, if token.starts_with('A') { "B" } else { "A" });
407
408 assert!(!ks.validates_gallery_media_token(&tampered, MEDIA_PATH, 1_900, 1_000));
409 assert!(!ks.validates_gallery_media_token(
410 &sign_gallery_media_token(&[0x24; GALLERY_SIGNING_SECRET_BYTES], MEDIA_PATH, 1_900,),
411 MEDIA_PATH,
412 1_900,
413 1_000,
414 ));
415 assert!(!ks.validates_gallery_media_token("not+url/safe", MEDIA_PATH, 1_900, 1_000,));
416 assert!(!ks.validates_gallery_media_token(
417 &token,
418 "/api/gallery/image/other.mp4",
419 1_900,
420 1_000,
421 ));
422 }
423
424 #[test]
425 fn gallery_media_token_cannot_be_reproduced_from_a_weak_api_key() {
426 const WEAK_API_KEY: &str = "password";
427 const MEDIA_PATH: &str = "/api/gallery/image/clip.mp4";
428 let ks = ApiKeySet::new_with_gallery_signing_secret(
429 HashSet::from([WEAK_API_KEY.to_string()]),
430 [0x42; GALLERY_SIGNING_SECRET_BYTES],
431 );
432 let token = ks.sign_gallery_media_token_for_tests(MEDIA_PATH, 1_900);
433 let api_key_derived = sign_gallery_media_token(WEAK_API_KEY.as_bytes(), MEDIA_PATH, 1_900);
434 let other_process =
435 sign_gallery_media_token(&[0x24; GALLERY_SIGNING_SECRET_BYTES], MEDIA_PATH, 1_900);
436
437 assert_ne!(token, api_key_derived);
438 assert_ne!(token, other_process);
439 assert!(!ks.validates_gallery_media_token(&api_key_derived, MEDIA_PATH, 1_900, 1_000,));
440 assert!(!ks.validates_gallery_media_token(&other_process, MEDIA_PATH, 1_900, 1_000,));
441 assert!(ks.contains(WEAK_API_KEY));
442 }
443}