1#![cfg(feature = "server")]
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub(crate) struct CompressionAlgorithms {
14 gzip: bool,
15 brotli: bool,
16 deflate: bool,
17 zstd: bool,
18}
19
20impl CompressionAlgorithms {
21 fn all_enabled() -> Self {
22 Self {
23 gzip: true,
24 brotli: true,
25 deflate: true,
26 zstd: true,
27 }
28 }
29
30 fn is_empty(&self) -> bool {
31 !self.gzip && !self.brotli && !self.deflate && !self.zstd
32 }
33}
34
35pub(crate) fn parse_compression_algorithms(env: &str) -> Option<CompressionAlgorithms> {
38 let env = env.trim();
39 if env.is_empty() || env.eq_ignore_ascii_case("none") || env.eq_ignore_ascii_case("off") {
40 return None;
41 }
42
43 let mut all = false;
44 let mut gzip = false;
45 let mut brotli = false;
46 let mut deflate = false;
47 let mut zstd = false;
48
49 for part in env.split(',') {
50 match part.trim().to_lowercase().as_str() {
51 "all" => all = true,
52 "gzip" => gzip = true,
53 "brotli" | "br" => brotli = true,
54 "deflate" => deflate = true,
55 "zstd" => zstd = true,
56 other => tracing::warn!(
57 "Unknown compression algorithm in COMPRESSION_ALGORITHMS: '{}'",
58 other
59 ),
60 }
61 }
62
63 if all {
64 return Some(CompressionAlgorithms::all_enabled());
65 }
66
67 let algorithms = CompressionAlgorithms {
68 gzip,
69 brotli,
70 deflate,
71 zstd,
72 };
73 if algorithms.is_empty() {
74 return None;
75 }
76
77 Some(algorithms)
78}
79
80pub(crate) fn compression_layer_from_env() -> Option<tower_http::compression::CompressionLayer> {
93 use tower_http::compression::CompressionLayer;
94
95 let env = std::env::var("COMPRESSION_ALGORITHMS").unwrap_or_else(|_| "off".to_string());
96 let algorithms = parse_compression_algorithms(&env)?;
97
98 Some(
99 CompressionLayer::new()
100 .gzip(algorithms.gzip)
101 .br(algorithms.brotli)
102 .deflate(algorithms.deflate)
103 .zstd(algorithms.zstd),
104 )
105}
106
107pub(crate) fn cache_control_for_path(
110 path: &str,
111 method: &axum::http::Method,
112) -> Option<axum::http::HeaderValue> {
113 use axum::http::{HeaderValue, Method};
114
115 if *method != Method::GET && *method != Method::HEAD {
117 return None;
118 }
119
120 if path == "/notes" || path.starts_with("/notes/") {
122 return Some(HeaderValue::from_static("no-store"));
123 }
124 if path.starts_with("/admin/notes")
125 || path == "/admin/notebooks"
126 || path.starts_with("/note-media/")
127 {
128 return Some(HeaderValue::from_static("private, no-store"));
129 }
130 if path.starts_with("/api") {
131 return None;
132 }
133
134 if path.starts_with("/admin") || path == "/login" || path == "/register" {
136 return None;
137 }
138
139 if path.starts_with("/_dioxus/")
142 || path.starts_with("/wasm/")
143 || path.ends_with(".wasm")
144 || path.ends_with(".js")
145 || path.ends_with(".css")
146 {
147 return Some(HeaderValue::from_static("public, no-cache"));
148 }
149
150 Some(HeaderValue::from_static(
152 "public, max-age=300, stale-while-revalidate=3600",
153 ))
154}
155
156pub(crate) async fn add_cache_control(
158 req: axum::extract::Request,
159 next: axum::middleware::Next,
160) -> axum::response::Response {
161 use axum::http::header;
162
163 let path = req.uri().path().to_string();
164 let cache_value = cache_control_for_path(&path, req.method());
165
166 let mut response = next.run(req).await;
167
168 if let Some(value) = cache_value {
169 response
171 .headers_mut()
172 .entry(header::CACHE_CONTROL)
173 .or_insert(value);
174 }
175
176 response
177}
178
179pub(crate) async fn admin_guard(
193 req: axum::extract::Request,
194 next: axum::middleware::Next,
195) -> axum::response::Response {
196 use crate::models::user::UserRole;
197 use axum::body::Body;
198 use axum::http::{header, StatusCode};
199 use axum::response::Response;
200
201 let path = req.uri().path().to_string();
202 if !path.starts_with("/admin") {
203 return next.run(req).await;
204 }
205
206 let cookie = req
208 .headers()
209 .get("cookie")
210 .and_then(|h| h.to_str().ok())
211 .unwrap_or("");
212 let token = crate::auth::session::parse_session_token(cookie);
213
214 let is_admin = match token {
215 Some(t) => match crate::api::auth::get_user_by_token(t).await {
216 Ok(Some(user)) => user.role == UserRole::Admin,
217 _ => true,
219 },
220 None => false,
222 };
223
224 if is_admin {
225 next.run(req).await
226 } else {
227 Response::builder()
228 .status(StatusCode::FOUND)
229 .header(header::LOCATION, "/login")
230 .body(Body::empty())
231 .expect("静态 302 重定向响应(合法 status + 固定 header + 空 body)必然构造成功")
232 }
233}
234
235pub(crate) async fn ssr_generation_middleware(
239 req: axum::extract::Request,
240 next: axum::middleware::Next,
241) -> axum::response::Response {
242 let generation = crate::ssr_cache::current_global_generation();
243 let is_get = req.method() == axum::http::Method::GET;
244 let (mut parts, body) = req.into_parts();
245 parts
246 .extensions
247 .insert(crate::ssr_cache::SsrGeneration(generation));
248 let mut response = next.run(axum::http::Request::from_parts(parts, body)).await;
249 if is_get {
250 response.headers_mut().insert(
251 axum::http::header::HeaderName::from_static("x-ssr-generation"),
252 axum::http::HeaderValue::from_str(&generation.to_string())
253 .unwrap_or_else(|_| axum::http::HeaderValue::from_static("0")),
254 );
255 }
256 response
257}
258
259pub(crate) async fn version_headers_middleware(
266 req: axum::extract::Request,
267 next: axum::middleware::Next,
268) -> axum::response::Response {
269 let mut response = next.run(req).await;
270 let h = response.headers_mut();
271 h.insert(
272 axum::http::header::SERVER,
273 axum::http::HeaderValue::from_str(&format!(
274 "yggdrasil/{}",
275 crate::build_info::BUILD_INFO.version
276 ))
277 .unwrap_or_else(|_| axum::http::HeaderValue::from_static("yggdrasil")),
278 );
279 h.insert(
280 axum::http::header::HeaderName::from_static("x-yggdrasil-version"),
281 axum::http::HeaderValue::from_static(crate::build_info::BUILD_INFO.version),
282 );
283 h.insert(
284 axum::http::header::HeaderName::from_static("x-yggdrasil-git"),
285 axum::http::HeaderValue::from_str(crate::build_info::BUILD_INFO.git_describe)
286 .unwrap_or_else(|_| axum::http::HeaderValue::from_static("unknown")),
287 );
288 h.insert(
289 axum::http::header::HeaderName::from_static("x-yggdrasil-hash"),
290 axum::http::HeaderValue::from_str(crate::build_info::BUILD_INFO.git_hash)
291 .unwrap_or_else(|_| axum::http::HeaderValue::from_static("unknown")),
292 );
293 response
294}
295
296#[cfg(test)]
297mod tests {
298 use super::{cache_control_for_path, parse_compression_algorithms, CompressionAlgorithms};
299 use axum::http::Method;
300
301 fn cache_value(path: &str, method: Method) -> Option<String> {
302 cache_control_for_path(path, &method).map(|v| v.to_str().unwrap().to_string())
303 }
304
305 #[test]
306 fn public_page_is_cached() {
307 assert_eq!(
308 cache_value("/", Method::GET),
309 Some("public, max-age=300, stale-while-revalidate=3600".to_string())
310 );
311 assert_eq!(
312 cache_value("/post/hello-world", Method::GET),
313 Some("public, max-age=300, stale-while-revalidate=3600".to_string())
314 );
315 assert_eq!(
316 cache_value("/tags/rust", Method::GET),
317 Some("public, max-age=300, stale-while-revalidate=3600".to_string())
318 );
319 }
320
321 #[test]
322 fn revocable_notes_never_use_browser_or_cdn_cache() {
323 for path in ["/notes", "/notes/example", "/notes/book/1"] {
324 for method in [Method::GET, Method::HEAD] {
325 assert_eq!(cache_value(path, method).as_deref(), Some("no-store"));
326 }
327 }
328 assert_eq!(
329 cache_value("/admin/notes/edit/1", Method::GET).as_deref(),
330 Some("private, no-store")
331 );
332 assert_eq!(
333 cache_value("/note-media/example", Method::GET).as_deref(),
334 Some("private, no-store")
335 );
336 }
337
338 #[test]
339 fn unversioned_assets_require_revalidation() {
340 for path in [
341 "/style.css",
342 "/highlight.css",
343 "/tiptap/editor.css",
344 "/xterm/terminal.css",
345 "/tiptap/editor.js",
346 "/codemirror/editor.js",
347 "/yggdrasil-core/yggdrasil-core.js",
348 "/mermaid/mermaid.js",
349 "/wasm/app.wasm",
350 "/wasm/app.js",
351 "/_dioxus/assets/main.js",
352 ] {
353 for method in [Method::GET, Method::HEAD] {
354 assert_eq!(
355 cache_value(path, method).as_deref(),
356 Some("public, no-cache"),
357 "{path}"
358 );
359 }
360 }
361 }
362
363 #[test]
364 fn api_and_admin_and_auth_are_not_cached() {
365 assert_eq!(cache_value("/api/posts", Method::GET), None);
366 assert_eq!(cache_value("/admin", Method::GET), None);
367 assert_eq!(cache_value("/admin/posts", Method::GET), None);
368 assert_eq!(cache_value("/login", Method::GET), None);
369 assert_eq!(cache_value("/register", Method::GET), None);
370 }
371
372 #[test]
373 fn non_get_requests_are_not_cached() {
374 assert_eq!(cache_value("/", Method::POST), None);
375 assert_eq!(cache_value("/post/hello-world", Method::POST), None);
376 assert_eq!(cache_value("/style.css", Method::POST), None);
377 }
378
379 #[test]
380 fn head_requests_are_cached_like_get() {
381 assert_eq!(
382 cache_value("/", Method::HEAD),
383 Some("public, max-age=300, stale-while-revalidate=3600".to_string())
384 );
385 }
386
387 #[test]
388 fn compression_all_enables_everything() {
389 assert_eq!(
390 parse_compression_algorithms("all"),
391 Some(CompressionAlgorithms::all_enabled())
392 );
393 }
394
395 #[test]
396 fn compression_default_env_is_off() {
397 assert_eq!(parse_compression_algorithms("off"), None);
399 }
400
401 #[test]
402 fn compression_empty_none_off_disable() {
403 assert_eq!(parse_compression_algorithms(""), None);
404 assert_eq!(parse_compression_algorithms("none"), None);
405 assert_eq!(parse_compression_algorithms("NONE"), None);
406 assert_eq!(parse_compression_algorithms("off"), None);
407 assert_eq!(parse_compression_algorithms("OFF"), None);
408 }
409
410 #[test]
411 fn compression_single_algorithm() {
412 assert_eq!(
413 parse_compression_algorithms("gzip"),
414 Some(CompressionAlgorithms {
415 gzip: true,
416 brotli: false,
417 deflate: false,
418 zstd: false,
419 })
420 );
421 assert_eq!(
422 parse_compression_algorithms("br"),
423 Some(CompressionAlgorithms {
424 gzip: false,
425 brotli: true,
426 deflate: false,
427 zstd: false,
428 })
429 );
430 }
431
432 #[test]
433 fn compression_multiple_algorithms() {
434 assert_eq!(
435 parse_compression_algorithms("gzip, zstd"),
436 Some(CompressionAlgorithms {
437 gzip: true,
438 brotli: false,
439 deflate: false,
440 zstd: true,
441 })
442 );
443 }
444
445 #[test]
446 fn compression_case_insensitive_and_whitespace_tolerant() {
447 assert_eq!(
448 parse_compression_algorithms("GZIP, Brotli, Deflate, Zstd"),
449 Some(CompressionAlgorithms::all_enabled())
450 );
451 assert_eq!(
452 parse_compression_algorithms(" gzip , br , deflate , zstd "),
453 Some(CompressionAlgorithms::all_enabled())
454 );
455 }
456
457 #[test]
458 fn compression_unknown_algorithms_are_ignored() {
459 assert_eq!(
460 parse_compression_algorithms("gzip, unknown, lz4"),
461 Some(CompressionAlgorithms {
462 gzip: true,
463 brotli: false,
464 deflate: false,
465 zstd: false,
466 })
467 );
468 }
469}