style: apply cargo fmt to entire codebase
Standardize code formatting across all 173 Rust source files using rustfmt. No functional changes - purely cosmetic. This establishes a consistent code style baseline for the project going forward.
This commit is contained in:
@@ -1,11 +1,11 @@
|
||||
use std::sync::Arc;
|
||||
use std::convert::Infallible;
|
||||
use axum::{
|
||||
extract::{State, Request, FromRequestParts},
|
||||
http::{StatusCode, HeaderMap, header, request::Parts},
|
||||
extract::{FromRequestParts, Request, State},
|
||||
http::{HeaderMap, StatusCode, header, request::Parts},
|
||||
middleware::Next,
|
||||
response::{Response, IntoResponse},
|
||||
response::{IntoResponse, Response},
|
||||
};
|
||||
use std::convert::Infallible;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::common::di::AppState;
|
||||
|
||||
@@ -77,7 +77,10 @@ where
|
||||
|
||||
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
|
||||
Ok(OptionalUserId(
|
||||
parts.extensions.get::<CurrentUser>().map(|cu| cu.id.clone()),
|
||||
parts
|
||||
.extensions
|
||||
.get::<CurrentUser>()
|
||||
.map(|cu| cu.id.clone()),
|
||||
))
|
||||
}
|
||||
}
|
||||
@@ -94,12 +97,12 @@ where
|
||||
type Rejection = Infallible;
|
||||
|
||||
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
|
||||
Ok(OptionalAuthUser(
|
||||
parts.extensions.get::<CurrentUser>().map(|cu| AuthUser {
|
||||
Ok(OptionalAuthUser(parts.extensions.get::<CurrentUser>().map(
|
||||
|cu| AuthUser {
|
||||
id: cu.id.clone(),
|
||||
username: cu.username.clone(),
|
||||
}),
|
||||
))
|
||||
},
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -108,19 +111,19 @@ where
|
||||
pub enum AuthError {
|
||||
#[error("Token not provided")]
|
||||
TokenNotProvided,
|
||||
|
||||
|
||||
#[error("Invalid token: {0}")]
|
||||
InvalidToken(String),
|
||||
|
||||
|
||||
#[error("Token expired")]
|
||||
TokenExpired,
|
||||
|
||||
|
||||
#[error("User not found")]
|
||||
UserNotFound,
|
||||
|
||||
|
||||
#[error("Access denied: {0}")]
|
||||
AccessDenied(String),
|
||||
|
||||
|
||||
#[error("Authentication service unavailable")]
|
||||
AuthServiceUnavailable,
|
||||
}
|
||||
@@ -128,12 +131,17 @@ pub enum AuthError {
|
||||
impl IntoResponse for AuthError {
|
||||
fn into_response(self) -> Response {
|
||||
let (status, error_message) = match self {
|
||||
AuthError::TokenNotProvided => (StatusCode::UNAUTHORIZED, "Token not provided".to_string()),
|
||||
AuthError::TokenNotProvided => {
|
||||
(StatusCode::UNAUTHORIZED, "Token not provided".to_string())
|
||||
}
|
||||
AuthError::InvalidToken(msg) => (StatusCode::UNAUTHORIZED, msg),
|
||||
AuthError::TokenExpired => (StatusCode::UNAUTHORIZED, "Token expired".to_string()),
|
||||
AuthError::UserNotFound => (StatusCode::UNAUTHORIZED, "User not found".to_string()),
|
||||
AuthError::AccessDenied(msg) => (StatusCode::FORBIDDEN, msg),
|
||||
AuthError::AuthServiceUnavailable => (StatusCode::INTERNAL_SERVER_ERROR, "Authentication service unavailable".to_string()),
|
||||
AuthError::AuthServiceUnavailable => (
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
"Authentication service unavailable".to_string(),
|
||||
),
|
||||
};
|
||||
|
||||
let body = axum::Json(serde_json::json!({
|
||||
@@ -160,15 +168,15 @@ pub async fn auth_middleware(
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| value.strip_prefix("Bearer "))
|
||||
.ok_or(AuthError::TokenNotProvided)?;
|
||||
|
||||
|
||||
// Validate that the token is not empty
|
||||
let token_str = token_str.trim();
|
||||
if token_str.is_empty() {
|
||||
return Err(AuthError::TokenNotProvided);
|
||||
}
|
||||
|
||||
|
||||
tracing::debug!("Processing authentication token");
|
||||
|
||||
|
||||
// Validate the token using the authentication service
|
||||
if let Some(auth_service) = state.auth_service.as_ref() {
|
||||
let token_service = &auth_service.token_service;
|
||||
@@ -183,14 +191,14 @@ pub async fn auth_middleware(
|
||||
};
|
||||
request.extensions_mut().insert(current_user);
|
||||
return Ok(next.run(request).await);
|
||||
},
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!("Token validation failed: {}", e);
|
||||
return Err(AuthError::InvalidToken(format!("Invalid token: {}", e)));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// If no authentication service is available, deny access
|
||||
tracing::error!("Auth middleware invoked but auth service is not configured");
|
||||
Err(AuthError::AuthServiceUnavailable)
|
||||
@@ -200,22 +208,23 @@ pub async fn auth_middleware(
|
||||
///
|
||||
/// Must be applied AFTER auth_middleware, as it depends on
|
||||
/// `CurrentUser` being present in the request extensions.
|
||||
pub async fn require_admin(
|
||||
request: Request,
|
||||
next: Next,
|
||||
) -> Response {
|
||||
pub async fn require_admin(request: Request, next: Next) -> Response {
|
||||
// Get the CurrentUser inserted by auth_middleware
|
||||
if let Some(current_user) = request.extensions().get::<CurrentUser>() {
|
||||
if current_user.role == "admin" {
|
||||
tracing::debug!("Admin access granted for user: {}", current_user.username);
|
||||
return next.run(request).await;
|
||||
}
|
||||
tracing::warn!("Admin access denied for user: {} (role: {})", current_user.username, current_user.role);
|
||||
tracing::warn!(
|
||||
"Admin access denied for user: {} (role: {})",
|
||||
current_user.username,
|
||||
current_user.role
|
||||
);
|
||||
} else {
|
||||
tracing::warn!("Admin check failed: no authenticated user in request");
|
||||
}
|
||||
|
||||
|
||||
// Access denied
|
||||
let error = AuthError::AccessDenied("Admin role required".to_string());
|
||||
error.into_response()
|
||||
}
|
||||
}
|
||||
|
||||
+178
-131
@@ -3,22 +3,22 @@ use axum::{
|
||||
http::{HeaderMap, HeaderValue, Method, Request, Response, StatusCode},
|
||||
middleware::Next,
|
||||
};
|
||||
use std::collections::hash_map::DefaultHasher;
|
||||
use std::hash::{Hash, Hasher};
|
||||
use std::time::{Duration, SystemTime};
|
||||
use bytes::Bytes;
|
||||
use chrono::{DateTime, Utc};
|
||||
use serde::Serialize;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::collections::HashMap;
|
||||
use tower::{Layer, Service};
|
||||
use std::task::{Context, Poll};
|
||||
use std::pin::Pin;
|
||||
use std::collections::hash_map::DefaultHasher;
|
||||
use std::future::Future;
|
||||
use bytes::Bytes;
|
||||
use std::hash::{Hash, Hasher};
|
||||
use std::pin::Pin;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::task::{Context, Poll};
|
||||
use std::time::{Duration, SystemTime};
|
||||
use tower::{Layer, Service};
|
||||
use tracing::{debug, info};
|
||||
|
||||
const MAX_CACHE_ENTRIES: usize = 1000; // Maximum number of cache entries
|
||||
const DEFAULT_MAX_AGE: u64 = 60; // Default time-to-live in seconds
|
||||
const MAX_CACHE_ENTRIES: usize = 1000; // Maximum number of cache entries
|
||||
const DEFAULT_MAX_AGE: u64 = 60; // Default time-to-live in seconds
|
||||
|
||||
// Type definitions for clarity
|
||||
type CacheKey = String;
|
||||
@@ -56,7 +56,7 @@ impl HttpCache {
|
||||
default_max_age: DEFAULT_MAX_AGE,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Creates a new instance with a specified time-to-live
|
||||
pub fn with_max_age(max_age: u64) -> Self {
|
||||
Self {
|
||||
@@ -64,65 +64,74 @@ impl HttpCache {
|
||||
default_max_age: max_age,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Gets cache statistics
|
||||
pub fn stats(&self) -> (usize, usize) {
|
||||
let lock = self.cache.lock().unwrap();
|
||||
let total = lock.len();
|
||||
|
||||
|
||||
// Count valid entries
|
||||
let _now = SystemTime::now();
|
||||
let valid = lock.values().filter(|entry| {
|
||||
match entry.timestamp.elapsed() {
|
||||
let valid = lock
|
||||
.values()
|
||||
.filter(|entry| match entry.timestamp.elapsed() {
|
||||
Ok(elapsed) => elapsed.as_secs() < entry.max_age,
|
||||
Err(_) => false,
|
||||
}
|
||||
}).count();
|
||||
|
||||
})
|
||||
.count();
|
||||
|
||||
(total, valid)
|
||||
}
|
||||
|
||||
|
||||
/// Cleans up expired entries
|
||||
pub fn cleanup(&self) -> usize {
|
||||
let mut lock = self.cache.lock().unwrap();
|
||||
let initial_count = lock.len();
|
||||
|
||||
|
||||
// Remove expired entries
|
||||
let _now = SystemTime::now();
|
||||
lock.retain(|_, entry| {
|
||||
match entry.timestamp.elapsed() {
|
||||
Ok(elapsed) => elapsed.as_secs() < entry.max_age,
|
||||
Err(_) => false,
|
||||
}
|
||||
lock.retain(|_, entry| match entry.timestamp.elapsed() {
|
||||
Ok(elapsed) => elapsed.as_secs() < entry.max_age,
|
||||
Err(_) => false,
|
||||
});
|
||||
|
||||
|
||||
let removed = initial_count - lock.len();
|
||||
debug!("HttpCache cleanup: removed {} expired entries", removed);
|
||||
|
||||
|
||||
removed
|
||||
}
|
||||
|
||||
|
||||
/// Sets an entry in the cache
|
||||
fn set(&self, key: &str, etag: EntityTag, data: Option<Bytes>, headers: HeaderMap, max_age: Option<u64>) {
|
||||
fn set(
|
||||
&self,
|
||||
key: &str,
|
||||
etag: EntityTag,
|
||||
data: Option<Bytes>,
|
||||
headers: HeaderMap,
|
||||
max_age: Option<u64>,
|
||||
) {
|
||||
let mut lock = self.cache.lock().unwrap();
|
||||
|
||||
|
||||
// Apply eviction policy if the cache is full
|
||||
if lock.len() >= MAX_CACHE_ENTRIES {
|
||||
debug!("Cache full, removing oldest entries");
|
||||
// Remove the oldest 10% of entries
|
||||
self.evict_oldest(&mut lock, MAX_CACHE_ENTRIES / 10);
|
||||
}
|
||||
|
||||
|
||||
// Store the new entry
|
||||
lock.insert(key.to_string(), CacheEntry {
|
||||
etag,
|
||||
data,
|
||||
headers,
|
||||
timestamp: SystemTime::now(),
|
||||
max_age: max_age.unwrap_or(self.default_max_age),
|
||||
});
|
||||
lock.insert(
|
||||
key.to_string(),
|
||||
CacheEntry {
|
||||
etag,
|
||||
data,
|
||||
headers,
|
||||
timestamp: SystemTime::now(),
|
||||
max_age: max_age.unwrap_or(self.default_max_age),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
/// Removes the oldest entries from the cache
|
||||
fn evict_oldest(&self, cache: &mut HashMap<CacheKey, CacheEntry>, count: usize) {
|
||||
// Sort by timestamp
|
||||
@@ -130,20 +139,20 @@ impl HttpCache {
|
||||
.iter()
|
||||
.map(|(key, entry)| (key.clone(), entry.timestamp))
|
||||
.collect();
|
||||
|
||||
|
||||
// Sort by timestamp (oldest first)
|
||||
entries.sort_by(|a, b| a.1.cmp(&b.1));
|
||||
|
||||
|
||||
// Remove the oldest entries
|
||||
for (key, _) in entries.iter().take(count) {
|
||||
cache.remove(key);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Gets an entry from the cache
|
||||
fn get(&self, key: &str) -> Option<CacheEntry> {
|
||||
let lock = self.cache.lock().unwrap();
|
||||
|
||||
|
||||
// Look up the entry
|
||||
if let Some(entry) = lock.get(key) {
|
||||
// Check if it has expired
|
||||
@@ -158,17 +167,17 @@ impl HttpCache {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
|
||||
/// Generates a simple ETag for a block of bytes
|
||||
fn calculate_etag_for_bytes(&self, bytes: &[u8]) -> EntityTag {
|
||||
// Calculate hash
|
||||
let mut hasher = DefaultHasher::new();
|
||||
bytes.hash(&mut hasher);
|
||||
let hash = hasher.finish();
|
||||
|
||||
|
||||
format!("\"{}\"", hash)
|
||||
}
|
||||
}
|
||||
@@ -180,80 +189,92 @@ pub async fn cache_middleware<T>(
|
||||
max_age: Option<u64>,
|
||||
req: Request<Body>,
|
||||
next: Next,
|
||||
) -> Result<Response<Body>, (StatusCode, String)>
|
||||
where
|
||||
T: Serialize
|
||||
) -> Result<Response<Body>, (StatusCode, String)>
|
||||
where
|
||||
T: Serialize,
|
||||
{
|
||||
// Only apply cache for GET requests
|
||||
if req.method() != Method::GET {
|
||||
return Ok(next.run(req).await);
|
||||
}
|
||||
|
||||
|
||||
// Check if the response is cached
|
||||
let if_none_match = req.headers()
|
||||
let if_none_match = req
|
||||
.headers()
|
||||
.get("if-none-match")
|
||||
.and_then(|v| v.to_str().ok());
|
||||
|
||||
|
||||
// If there is a cache entry
|
||||
if let Some(cache_entry) = cache.get(cache_key) {
|
||||
// Check if the client already has the updated version
|
||||
if let Some(client_etag) = if_none_match
|
||||
&& client_etag == cache_entry.etag {
|
||||
// The client has the most recent version, send 304 Not Modified
|
||||
debug!("Cache hit (304) for key: {}", cache_key);
|
||||
return Ok(create_not_modified_response(&cache_entry));
|
||||
}
|
||||
|
||||
&& client_etag == cache_entry.etag
|
||||
{
|
||||
// The client has the most recent version, send 304 Not Modified
|
||||
debug!("Cache hit (304) for key: {}", cache_key);
|
||||
return Ok(create_not_modified_response(&cache_entry));
|
||||
}
|
||||
|
||||
// The client needs the updated version
|
||||
if let Some(data) = &cache_entry.data {
|
||||
debug!("Cache hit (200) for key: {}", cache_key);
|
||||
|
||||
|
||||
// Create response with cached data
|
||||
let mut response = Response::new(Body::from(data.clone()));
|
||||
|
||||
|
||||
// Copy original headers
|
||||
for (key, value) in &cache_entry.headers {
|
||||
if !key.as_str().eq_ignore_ascii_case("transfer-encoding") {
|
||||
response.headers_mut().insert(key.clone(), value.clone());
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Add cache headers
|
||||
set_cache_headers(&mut response, &cache_entry.etag, max_age.unwrap_or(cache_entry.max_age));
|
||||
|
||||
set_cache_headers(
|
||||
&mut response,
|
||||
&cache_entry.etag,
|
||||
max_age.unwrap_or(cache_entry.max_age),
|
||||
);
|
||||
|
||||
return Ok(response);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Not cached or expired, continue with the middleware
|
||||
debug!("Cache miss for key: {}", cache_key);
|
||||
let response = next.run(req).await;
|
||||
|
||||
|
||||
// Don't cache errors
|
||||
if !response.status().is_success() {
|
||||
return Ok(response);
|
||||
}
|
||||
|
||||
|
||||
// Convert the response to calculate the ETag
|
||||
let (parts, _body) = response.into_parts();
|
||||
let bytes = axum::body::to_bytes(_body, 1024 * 1024 * 10).await.unwrap_or_default();
|
||||
|
||||
let bytes = axum::body::to_bytes(_body, 1024 * 1024 * 10)
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
|
||||
// Calculate ETag
|
||||
let etag = cache.calculate_etag_for_bytes(&bytes);
|
||||
|
||||
|
||||
// Save to cache
|
||||
cache.set(
|
||||
cache_key,
|
||||
etag.clone(),
|
||||
cache_key,
|
||||
etag.clone(),
|
||||
Some(bytes.clone()),
|
||||
parts.headers.clone(),
|
||||
max_age
|
||||
max_age,
|
||||
);
|
||||
|
||||
|
||||
// Create the response with ETag
|
||||
let mut response = Response::from_parts(parts, Body::from(bytes));
|
||||
set_cache_headers(&mut response, &etag, max_age.unwrap_or(cache.default_max_age));
|
||||
|
||||
set_cache_headers(
|
||||
&mut response,
|
||||
&etag,
|
||||
max_age.unwrap_or(cache.default_max_age),
|
||||
);
|
||||
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
@@ -263,18 +284,20 @@ fn create_not_modified_response(entry: &CacheEntry) -> Response<Body> {
|
||||
.status(StatusCode::NOT_MODIFIED)
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
|
||||
|
||||
// Copy cache headers
|
||||
if let Some(cache_control) = entry.headers.get("cache-control") {
|
||||
response.headers_mut().insert("cache-control", cache_control.clone());
|
||||
response
|
||||
.headers_mut()
|
||||
.insert("cache-control", cache_control.clone());
|
||||
}
|
||||
|
||||
|
||||
// Add ETag
|
||||
response.headers_mut().insert(
|
||||
"etag",
|
||||
HeaderValue::from_str(&entry.etag).unwrap_or(HeaderValue::from_static(""))
|
||||
"etag",
|
||||
HeaderValue::from_str(&entry.etag).unwrap_or(HeaderValue::from_static("")),
|
||||
);
|
||||
|
||||
|
||||
response
|
||||
}
|
||||
|
||||
@@ -282,23 +305,23 @@ fn create_not_modified_response(entry: &CacheEntry) -> Response<Body> {
|
||||
fn set_cache_headers(response: &mut Response<Body>, etag: &str, max_age: u64) {
|
||||
// Add ETag
|
||||
response.headers_mut().insert(
|
||||
"etag",
|
||||
HeaderValue::from_str(etag).unwrap_or(HeaderValue::from_static(""))
|
||||
"etag",
|
||||
HeaderValue::from_str(etag).unwrap_or(HeaderValue::from_static("")),
|
||||
);
|
||||
|
||||
|
||||
// Configure Cache-Control
|
||||
let cache_control = format!("public, max-age={}", max_age);
|
||||
response.headers_mut().insert(
|
||||
"cache-control",
|
||||
HeaderValue::from_str(&cache_control).unwrap_or(HeaderValue::from_static(""))
|
||||
HeaderValue::from_str(&cache_control).unwrap_or(HeaderValue::from_static("")),
|
||||
);
|
||||
|
||||
|
||||
// Add Last-Modified header
|
||||
let now: DateTime<Utc> = Utc::now();
|
||||
let last_modified = now.format("%a, %d %b %Y %H:%M:%S GMT").to_string();
|
||||
response.headers_mut().insert(
|
||||
"last-modified",
|
||||
HeaderValue::from_str(&last_modified).unwrap_or(HeaderValue::from_static(""))
|
||||
HeaderValue::from_str(&last_modified).unwrap_or(HeaderValue::from_static("")),
|
||||
);
|
||||
}
|
||||
|
||||
@@ -317,7 +340,7 @@ impl HttpCacheLayer {
|
||||
max_age: None,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Sets the maximum time-to-live
|
||||
pub fn with_max_age(mut self, max_age: u64) -> Self {
|
||||
self.max_age = Some(max_age);
|
||||
@@ -327,7 +350,7 @@ impl HttpCacheLayer {
|
||||
|
||||
impl<S> Layer<S> for HttpCacheLayer {
|
||||
type Service = HttpCacheService<S>;
|
||||
|
||||
|
||||
fn layer(&self, service: S) -> Self::Service {
|
||||
HttpCacheService {
|
||||
inner: service,
|
||||
@@ -358,15 +381,15 @@ where
|
||||
type Response = Response<Body>;
|
||||
type Error = Box<dyn std::error::Error + Send + Sync>;
|
||||
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
|
||||
|
||||
|
||||
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
||||
self.inner.poll_ready(cx).map_err(|e| e.into())
|
||||
}
|
||||
|
||||
|
||||
fn call(&mut self, req: Request<ReqBody>) -> Self::Future {
|
||||
// Generate cache key
|
||||
let cache_key = req.uri().path().to_string();
|
||||
|
||||
|
||||
// Only apply cache for GET requests
|
||||
if req.method() != Method::GET {
|
||||
let future = self.inner.call(req);
|
||||
@@ -375,41 +398,46 @@ where
|
||||
Ok(response_map_body(response).await)
|
||||
});
|
||||
}
|
||||
|
||||
|
||||
// Get client ETag
|
||||
let if_none_match = req.headers()
|
||||
let if_none_match = req
|
||||
.headers()
|
||||
.get("if-none-match")
|
||||
.and_then(|v| v.to_str().ok());
|
||||
|
||||
|
||||
// Check if there is a cache entry
|
||||
let cache_clone = self.cache.clone();
|
||||
let max_age = self.max_age;
|
||||
let entry = cache_clone.get(&cache_key);
|
||||
|
||||
|
||||
match entry {
|
||||
Some(cache_entry) if if_none_match == Some(&cache_entry.etag) => {
|
||||
// The client has the correct version, send 304
|
||||
debug!("Cache HIT (304): {}", cache_key);
|
||||
let response = create_not_modified_response(&cache_entry);
|
||||
Box::pin(async move { Ok(response) })
|
||||
},
|
||||
}
|
||||
Some(cache_entry) if cache_entry.data.is_some() => {
|
||||
// The client needs the updated version
|
||||
debug!("Cache HIT (200): {}", cache_key);
|
||||
let mut response = Response::new(Body::from(cache_entry.data.clone().unwrap()));
|
||||
|
||||
|
||||
// Copy original headers
|
||||
for (key, value) in &cache_entry.headers {
|
||||
if !key.as_str().eq_ignore_ascii_case("transfer-encoding") {
|
||||
response.headers_mut().insert(key.clone(), value.clone());
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Add cache headers
|
||||
set_cache_headers(&mut response, &cache_entry.etag, max_age.unwrap_or(cache_entry.max_age));
|
||||
|
||||
set_cache_headers(
|
||||
&mut response,
|
||||
&cache_entry.etag,
|
||||
max_age.unwrap_or(cache_entry.max_age),
|
||||
);
|
||||
|
||||
Box::pin(async move { Ok(response) })
|
||||
},
|
||||
}
|
||||
_ => {
|
||||
// Not cached or expired
|
||||
debug!("Cache MISS: {}", cache_key);
|
||||
@@ -417,36 +445,40 @@ where
|
||||
let cache_clone = self.cache.clone();
|
||||
let max_age = self.max_age;
|
||||
let cache_key = cache_key.clone();
|
||||
|
||||
|
||||
Box::pin(async move {
|
||||
let response = future.await.map_err(|e| e.into())?;
|
||||
let response = response_map_body(response).await;
|
||||
|
||||
|
||||
// Don't cache errors
|
||||
if !response.status().is_success() {
|
||||
return Ok(response);
|
||||
}
|
||||
|
||||
|
||||
// Get the body and calculate ETag
|
||||
let (parts, body) = response.into_parts();
|
||||
let bytes = axum::body::to_bytes(body, 1024 * 1024 * 10).await?;
|
||||
|
||||
|
||||
// Calculate ETag
|
||||
let etag = cache_clone.calculate_etag_for_bytes(&bytes);
|
||||
|
||||
|
||||
// Save to cache
|
||||
cache_clone.set(
|
||||
&cache_key,
|
||||
etag.clone(),
|
||||
&cache_key,
|
||||
etag.clone(),
|
||||
Some(bytes.clone()),
|
||||
parts.headers.clone(),
|
||||
max_age
|
||||
max_age,
|
||||
);
|
||||
|
||||
|
||||
// Create the response with ETag
|
||||
let mut response = Response::from_parts(parts, Body::from(bytes));
|
||||
set_cache_headers(&mut response, &etag, max_age.unwrap_or(cache_clone.default_max_age));
|
||||
|
||||
set_cache_headers(
|
||||
&mut response,
|
||||
&etag,
|
||||
max_age.unwrap_or(cache_clone.default_max_age),
|
||||
);
|
||||
|
||||
Ok(response)
|
||||
})
|
||||
}
|
||||
@@ -481,13 +513,16 @@ where
|
||||
pub fn start_cache_cleanup_task(cache: HttpCache) {
|
||||
tokio::spawn(async move {
|
||||
let mut interval = tokio::time::interval(Duration::from_secs(300)); // Every 5 minutes
|
||||
|
||||
|
||||
loop {
|
||||
interval.tick().await;
|
||||
let removed = cache.cleanup();
|
||||
let (total, valid) = cache.stats();
|
||||
|
||||
info!("HTTP Cache cleanup: removed {}, current: {}/{}", removed, valid, total);
|
||||
|
||||
info!(
|
||||
"HTTP Cache cleanup: removed {}, current: {}/{}",
|
||||
removed, valid, total
|
||||
);
|
||||
}
|
||||
});
|
||||
}
|
||||
@@ -496,49 +531,61 @@ pub fn start_cache_cleanup_task(cache: HttpCache) {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, Hash)]
|
||||
struct TestData {
|
||||
id: u32,
|
||||
name: String,
|
||||
}
|
||||
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_etag_generation() {
|
||||
let cache = HttpCache::new();
|
||||
|
||||
let data1 = serde_json::to_vec(&TestData { id: 1, name: "Test".to_string() }).unwrap();
|
||||
let data2 = serde_json::to_vec(&TestData { id: 1, name: "Test".to_string() }).unwrap();
|
||||
let data3 = serde_json::to_vec(&TestData { id: 2, name: "Test".to_string() }).unwrap();
|
||||
|
||||
|
||||
let data1 = serde_json::to_vec(&TestData {
|
||||
id: 1,
|
||||
name: "Test".to_string(),
|
||||
})
|
||||
.unwrap();
|
||||
let data2 = serde_json::to_vec(&TestData {
|
||||
id: 1,
|
||||
name: "Test".to_string(),
|
||||
})
|
||||
.unwrap();
|
||||
let data3 = serde_json::to_vec(&TestData {
|
||||
id: 2,
|
||||
name: "Test".to_string(),
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let etag1 = cache.calculate_etag_for_bytes(&data1);
|
||||
let etag2 = cache.calculate_etag_for_bytes(&data2);
|
||||
let etag3 = cache.calculate_etag_for_bytes(&data3);
|
||||
|
||||
|
||||
// Same data should generate the same ETag
|
||||
assert_eq!(etag1, etag2);
|
||||
|
||||
|
||||
// Different data should generate different ETags
|
||||
assert_ne!(etag1, etag3);
|
||||
}
|
||||
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_cache_hit_miss() {
|
||||
let cache = HttpCache::new();
|
||||
|
||||
|
||||
// Create test data directly as Bytes
|
||||
let bytes1 = Bytes::from(r#"{"id":1,"name":"Test"}"#);
|
||||
let headers1 = HeaderMap::new();
|
||||
|
||||
|
||||
let etag1 = cache.calculate_etag_for_bytes(&bytes1);
|
||||
cache.set("test", etag1.clone(), Some(bytes1.clone()), headers1, None);
|
||||
|
||||
|
||||
// Verify cache hit
|
||||
let entry = cache.get("test").unwrap();
|
||||
assert_eq!(entry.etag, etag1);
|
||||
assert_eq!(entry.data.unwrap(), bytes1);
|
||||
|
||||
|
||||
// Verify cache miss
|
||||
assert!(cache.get("nonexistent").is_none());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
pub mod cache;
|
||||
pub mod auth;
|
||||
pub mod redirect; // Add redirect middleware for API to Axum transition
|
||||
pub mod cache;
|
||||
pub mod redirect; // Add redirect middleware for API to Axum transition
|
||||
|
||||
@@ -1,12 +1,8 @@
|
||||
use std::task::{Context, Poll};
|
||||
use axum::http::{Uri, uri::PathAndQuery};
|
||||
use axum::{extract::Request, middleware::Next, response::Response};
|
||||
use std::future::Future;
|
||||
use std::pin::Pin;
|
||||
use axum::{
|
||||
extract::Request,
|
||||
response::Response,
|
||||
middleware::Next,
|
||||
};
|
||||
use axum::http::{uri::PathAndQuery, Uri};
|
||||
use std::task::{Context, Poll};
|
||||
use tower::{Layer, Service};
|
||||
|
||||
/// A middleware that redirects specific paths to the proper Axum routes.
|
||||
@@ -15,7 +11,7 @@ pub struct RedirectMiddleware<S> {
|
||||
inner: S,
|
||||
}
|
||||
|
||||
impl<S> Service<Request> for RedirectMiddleware<S>
|
||||
impl<S> Service<Request> for RedirectMiddleware<S>
|
||||
where
|
||||
S: Service<Request, Response = Response> + Send + 'static,
|
||||
S::Future: Send + 'static,
|
||||
@@ -33,7 +29,7 @@ where
|
||||
// Log the incoming request
|
||||
let uri = request.uri().clone();
|
||||
let path = uri.path().to_string();
|
||||
|
||||
|
||||
// Check and potentially redirect file-related API routes
|
||||
if path.starts_with("/api/files") {
|
||||
// Handle file-related redirects
|
||||
@@ -44,23 +40,28 @@ where
|
||||
// File download request - let's adjust the URI to match the Axum route
|
||||
// Extract the ID from the path
|
||||
let file_id = &path[11..];
|
||||
tracing::info!("Redirecting file download request: {} to /api/files/{}", path, file_id);
|
||||
|
||||
tracing::info!(
|
||||
"Redirecting file download request: {} to /api/files/{}",
|
||||
path,
|
||||
file_id
|
||||
);
|
||||
|
||||
// Create a new URI for the Axum route
|
||||
let uri_clone = uri.clone();
|
||||
let mut parts = uri_clone.into_parts();
|
||||
let query = parts.path_and_query
|
||||
let query = parts
|
||||
.path_and_query
|
||||
.as_ref()
|
||||
.and_then(|pq| pq.query())
|
||||
.map(|q| format!("?{}", q))
|
||||
.unwrap_or_default();
|
||||
|
||||
|
||||
let new_path = format!("/api/files/{}{}", file_id, query);
|
||||
parts.path_and_query = Some(
|
||||
PathAndQuery::from_maybe_shared(new_path.into_bytes())
|
||||
.expect("Failed to create path and query")
|
||||
.expect("Failed to create path and query"),
|
||||
);
|
||||
|
||||
|
||||
let new_uri = Uri::from_parts(parts).expect("Failed to create URI");
|
||||
*request.uri_mut() = new_uri;
|
||||
}
|
||||
@@ -69,10 +70,10 @@ where
|
||||
tracing::debug!("Folder request detected: {}", path);
|
||||
// We might need to add specific redirects for folder operations here
|
||||
}
|
||||
|
||||
|
||||
// Pass the request to the inner service
|
||||
let future = self.inner.call(request);
|
||||
|
||||
|
||||
Box::pin(async move {
|
||||
let response = future.await?;
|
||||
Ok(response)
|
||||
@@ -93,27 +94,27 @@ impl<S> Layer<S> for RedirectLayer {
|
||||
}
|
||||
|
||||
/// Axum middleware function that can be applied directly to routes
|
||||
pub async fn redirect_middleware(
|
||||
request: Request,
|
||||
next: Next,
|
||||
) -> Response {
|
||||
pub async fn redirect_middleware(request: Request, next: Next) -> Response {
|
||||
// Get the path
|
||||
let path = request.uri().path().to_string();
|
||||
|
||||
|
||||
// Process the request based on the path
|
||||
if path.starts_with("/api/files") || path.starts_with("/api/folders") || path.starts_with("/api/auth") {
|
||||
if path.starts_with("/api/files")
|
||||
|| path.starts_with("/api/folders")
|
||||
|| path.starts_with("/api/auth")
|
||||
{
|
||||
tracing::debug!("API request detected in middleware: {}", path);
|
||||
// Log additional information about the request
|
||||
if let Some(content_type) = request.headers().get("content-type") {
|
||||
tracing::debug!("Content-Type: {:?}", content_type);
|
||||
}
|
||||
|
||||
|
||||
// For debugging auth-related requests
|
||||
if path.starts_with("/api/auth") {
|
||||
tracing::info!("Auth API request: {} method: {}", path, request.method());
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Continue the middleware chain
|
||||
next.run(request).await
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user