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:
Dionisio
2026-02-14 01:29:34 +01:00
parent 67137a3ef2
commit 4c98c5a657
173 changed files with 23368 additions and 17590 deletions
+102 -95
View File
@@ -1,8 +1,8 @@
use std::cmp::min;
use std::collections::VecDeque;
use std::sync::Arc;
use tokio::sync::{Mutex, Semaphore};
use std::time::{Duration, Instant};
use tokio::sync::{Mutex, Semaphore};
use tracing::debug;
/// Default buffer size in the pool
@@ -79,16 +79,12 @@ impl BufferPool {
buffer_ttl: Duration::from_secs(buffer_ttl_secs),
})
}
/// Creates a pool with default configuration
pub fn default() -> Arc<Self> {
Self::new(
DEFAULT_BUFFER_SIZE,
DEFAULT_MAX_BUFFERS,
DEFAULT_BUFFER_TTL
)
Self::new(DEFAULT_BUFFER_SIZE, DEFAULT_MAX_BUFFERS, DEFAULT_BUFFER_TTL)
}
/// Gets a buffer from the pool or creates a new one if needed.
/// This version takes an Arc<Self> to ensure the BorrowedBuffer keeps a proper
/// reference to the shared pool (not a clone).
@@ -99,7 +95,7 @@ impl BufferPool {
let mut stats = self.stats.lock().await;
stats.gets += 1;
}
// Concurrency control
// Acquire a semaphore permit. If none available, wait.
// We forget() the permit so it doesn't auto-release on drop.
@@ -113,19 +109,23 @@ impl BufferPool {
stats.waits += 1;
stats.max_buffers_reached += 1;
}
debug!("Buffer pool: waiting for available buffer");
let permit = self.limit.acquire().await.expect("Semaphore should not be closed");
let permit = self
.limit
.acquire()
.await
.expect("Semaphore should not be closed");
debug!("Buffer pool: acquired buffer after waiting");
permit.forget();
}
};
// Try to get an existing buffer from the pool
let mut pool_locked = self.pool.lock().await;
let pool_arc = Arc::clone(self);
if let Some(mut pooled_buffer) = pool_locked.pop_front() {
// Check if the buffer has expired
if pooled_buffer.last_used.elapsed() > self.buffer_ttl {
@@ -134,12 +134,12 @@ impl BufferPool {
stats.evictions += 1;
stats.misses += 1;
drop(stats);
debug!("Buffer pool: evicted expired buffer");
// Create new buffer (reusing the permit)
drop(pool_locked); // Release the lock before returning
BorrowedBuffer {
buffer: vec![0; self.buffer_size],
used_size: 0,
@@ -151,13 +151,13 @@ impl BufferPool {
let mut stats = self.stats.lock().await;
stats.hits += 1;
drop(stats);
// Release the lock before returning
drop(pool_locked);
// Clear buffer for security
pooled_buffer.buffer.fill(0);
BorrowedBuffer {
buffer: pooled_buffer.buffer,
used_size: 0,
@@ -170,12 +170,12 @@ impl BufferPool {
let mut stats = self.stats.lock().await;
stats.misses += 1;
drop(stats);
// Release the lock before returning
drop(pool_locked);
debug!("Buffer pool: creating new buffer");
BorrowedBuffer {
buffer: vec![0; self.buffer_size],
used_size: 0,
@@ -184,90 +184,97 @@ impl BufferPool {
}
}
}
/// Returns a buffer to the pool
async fn return_buffer(&self, mut buffer: Vec<u8>) {
// If the buffer is the wrong size, discard it
if buffer.capacity() != self.buffer_size {
debug!("Buffer pool: discarding buffer of wrong size: {} (expected {})",
buffer.capacity(), self.buffer_size);
debug!(
"Buffer pool: discarding buffer of wrong size: {} (expected {})",
buffer.capacity(),
self.buffer_size
);
// Release the semaphore permit even if we discard the buffer
self.limit.add_permits(1);
return;
}
// Resize to ensure correct capacity
buffer.resize(self.buffer_size, 0);
// Add to the pool
let mut pool_locked = self.pool.lock().await;
pool_locked.push_back(PooledBuffer {
buffer,
last_used: Instant::now(),
});
// Update statistics
let mut stats = self.stats.lock().await;
stats.returns += 1;
// Release the semaphore permit so another caller can acquire a buffer
drop(pool_locked);
drop(stats);
self.limit.add_permits(1);
}
/// Cleans expired buffers from the pool
pub async fn clean_expired_buffers(&self) {
let _now = Instant::now();
let mut pool_locked = self.pool.lock().await;
// Count expired
let count_before = pool_locked.len();
// Filter keeping only non-expired
pool_locked.retain(|buffer| {
buffer.last_used.elapsed() <= self.buffer_ttl
});
pool_locked.retain(|buffer| buffer.last_used.elapsed() <= self.buffer_ttl);
// Count how many were removed
let removed = count_before - pool_locked.len();
if removed > 0 {
// Update statistics
let mut stats = self.stats.lock().await;
stats.evictions += removed;
debug!("Buffer pool: cleaned {} expired buffers", removed);
}
}
/// Gets current pool statistics
pub async fn get_stats(&self) -> BufferPoolStats {
self.stats.lock().await.clone()
}
/// Starts the periodic cleanup task
pub fn start_cleaner(pool: Arc<Self>) {
tokio::spawn(async move {
let interval = Duration::from_secs(30); // Clean every 30 seconds
loop {
tokio::time::sleep(interval).await;
pool.clean_expired_buffers().await;
// Log statistics periodically
let stats = pool.get_stats().await;
debug!("Buffer pool stats: gets={}, hits={}, misses={}, hit_ratio={:.2}%, returns={}, \
debug!(
"Buffer pool stats: gets={}, hits={}, misses={}, hit_ratio={:.2}%, returns={}, \
evictions={}, max_reached={}, waits={}",
stats.gets,
stats.hits,
stats.misses,
if stats.gets > 0 { (stats.hits as f64 * 100.0) / stats.gets as f64 } else { 0.0 },
stats.returns,
stats.evictions,
stats.max_buffers_reached,
stats.waits);
stats.gets,
stats.hits,
stats.misses,
if stats.gets > 0 {
(stats.hits as f64 * 100.0) / stats.gets as f64
} else {
0.0
},
stats.returns,
stats.evictions,
stats.max_buffers_reached,
stats.waits
);
}
});
}
@@ -290,26 +297,26 @@ impl BorrowedBuffer {
pub fn as_mut_slice(&mut self) -> &mut [u8] {
&mut self.buffer
}
/// Gets a reference to the used data
pub fn as_slice(&self) -> &[u8] {
&self.buffer[..self.used_size]
}
/// Sets how many bytes were actually used
pub fn set_used(&mut self, size: usize) {
self.used_size = min(size, self.buffer.len());
}
/// Converts into a Vec<u8> that includes only the used data
pub fn into_vec(mut self) -> Vec<u8> {
// Mark to not return to pool
self.return_to_pool = false;
// Create a new vector with only the used data
self.buffer[..self.used_size].to_vec()
}
/// Copies data to this buffer and updates the used size
pub fn copy_from_slice(&mut self, data: &[u8]) -> usize {
let copy_size = min(data.len(), self.buffer.len());
@@ -317,18 +324,18 @@ impl BorrowedBuffer {
self.used_size = copy_size;
copy_size
}
/// Prevents the buffer from being returned to the pool on destruction
pub fn do_not_return(mut self) -> Self {
self.return_to_pool = false;
self
}
/// Gets the total buffer size
pub fn capacity(&self) -> usize {
self.buffer.len()
}
/// Gets the used buffer size
pub fn used_size(&self) -> usize {
self.used_size
@@ -342,7 +349,7 @@ impl Drop for BorrowedBuffer {
// Take ownership of the buffer and create a clone of the pool
let buffer = std::mem::take(&mut self.buffer);
let pool = self.pool.clone();
// Spawn the return so that drop doesn't block
// return_buffer will release the semaphore permit
tokio::spawn(async move {
@@ -358,140 +365,140 @@ impl Drop for BorrowedBuffer {
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_buffer_pooling() {
// Create small pool for testing
let pool = BufferPool::new(1024, 5, 60);
// Get a buffer
let mut buffer1 = pool.get_buffer().await;
buffer1.copy_from_slice(b"test data");
assert_eq!(buffer1.as_slice(), b"test data");
// Get another buffer
let buffer2 = pool.get_buffer().await;
// Verify stats
let stats = pool.get_stats().await;
assert_eq!(stats.gets, 2);
assert_eq!(stats.hits, 0); // no hits yet
assert_eq!(stats.misses, 2); // all are misses
// Return buffer1 to pool (implicitly via drop)
drop(buffer1);
// Allow the async return to occur
tokio::time::sleep(Duration::from_millis(10)).await;
// Get another buffer (should reuse the returned one)
let buffer3 = pool.get_buffer().await;
// Verify updated stats
let stats = pool.get_stats().await;
assert_eq!(stats.gets, 3);
assert_eq!(stats.hits, 1); // now there should be a hit
assert_eq!(stats.returns, 1); // one buffer returned
// Cleanup
drop(buffer2);
drop(buffer3);
}
#[tokio::test]
async fn test_buffer_operations() {
let pool = BufferPool::new(1024, 10, 60);
// Get buffer
let mut buffer = pool.get_buffer().await;
// Write data
buffer.copy_from_slice(b"Hello, world!");
assert_eq!(buffer.used_size(), 13);
assert_eq!(buffer.as_slice(), b"Hello, world!");
// Convert to vec and verify
let vec = buffer.into_vec(); // This prevents returning to pool
assert_eq!(vec, b"Hello, world!");
// Verify that returns are not incremented (buffer not returned)
tokio::time::sleep(Duration::from_millis(10)).await;
let stats = pool.get_stats().await;
assert_eq!(stats.returns, 0);
}
#[tokio::test]
async fn test_pool_limit() {
// Pool with only 3 buffers
let pool = BufferPool::new(1024, 3, 60);
// Get 3 buffers (reaches the limit)
let buffer1 = pool.get_buffer().await;
let buffer2 = pool.get_buffer().await;
let buffer3 = pool.get_buffer().await;
// Verify stats
let stats = pool.get_stats().await;
assert_eq!(stats.gets, 3);
assert_eq!(stats.waits, 0); // no waits yet
// Try to get a 4th buffer in a separate task (should wait)
let pool_clone = pool.clone();
let handle = tokio::spawn(async move {
let _buffer4 = pool_clone.get_buffer().await;
true
});
// Give time for the task to try to take the buffer
tokio::time::sleep(Duration::from_millis(50)).await;
// Verify there is a wait
let stats = pool.get_stats().await;
assert_eq!(stats.waits, 1);
// Release a buffer
drop(buffer1);
// Give time for the async return and for the waiting task to get its buffer
tokio::time::sleep(Duration::from_millis(50)).await;
// Verify the task was able to continue
assert!(handle.await.unwrap());
// Cleanup
drop(buffer2);
drop(buffer3);
}
#[tokio::test]
async fn test_ttl_expiration() {
// Pool with very short TTL for testing
let pool = BufferPool::new(1024, 5, 1); // 1 second TTL
// Get and return a buffer
let buffer = pool.get_buffer().await;
drop(buffer);
// Allow the async return to occur
tokio::time::sleep(Duration::from_millis(50)).await;
// Verify there is a buffer in the pool
let stats = pool.get_stats().await;
assert_eq!(stats.returns, 1);
// Wait for the TTL to expire
tokio::time::sleep(Duration::from_secs(2)).await;
// Clean expired
pool.clean_expired_buffers().await;
// Get another buffer (should be a miss since the previous one expired)
let _buffer2 = pool.get_buffer().await;
// Verify stats
let stats = pool.get_stats().await;
assert_eq!(stats.evictions, 1); // one expired buffer
assert_eq!(stats.hits, 0); // no hits (the buffer expired)
assert_eq!(stats.misses, 2); // two misses (1st and 3rd get)
}
}
}
File diff suppressed because it is too large Load Diff
+160 -100
View File
@@ -1,17 +1,16 @@
use std::io::{Read};
use std::sync::Arc;
use async_trait::async_trait;
use bytes::Bytes;
use futures::{Stream, StreamExt};
use tracing::error;
use std::io;
use flate2::Compression;
use flate2::read::GzEncoder as GzEncoderRead;
use flate2::bufread::GzDecoder;
use flate2::read::GzEncoder as GzEncoderRead;
use futures::{Stream, StreamExt};
use std::io;
use std::io::Read;
use std::sync::Arc;
use tracing::error;
use crate::application::ports::compression_ports::{
CompressionPort,
CompressionLevel as PortCompressionLevel,
CompressionLevel as PortCompressionLevel, CompressionPort,
};
use crate::domain::errors::DomainError;
use crate::infrastructure::services::buffer_pool::BufferPool;
@@ -48,22 +47,27 @@ const COMPRESSION_SIZE_THRESHOLD: u64 = 1024 * 50; // 50KB
pub trait CompressionService: Send + Sync {
/// Compresses data in memory
async fn compress_data(&self, data: &[u8], level: CompressionLevel) -> io::Result<Vec<u8>>;
/// Decompresses data in memory
async fn decompress_data(&self, compressed_data: &[u8]) -> io::Result<Vec<u8>>;
/// Compresses a data stream
fn compress_stream<S>(&self, stream: S, level: CompressionLevel)
-> impl Stream<Item = io::Result<Bytes>> + Send
fn compress_stream<S>(
&self,
stream: S,
level: CompressionLevel,
) -> impl Stream<Item = io::Result<Bytes>> + Send
where
S: Stream<Item = io::Result<Bytes>> + Send + 'static + Unpin;
/// Decompresses a data stream
fn decompress_stream<S>(&self, compressed_stream: S)
-> impl Stream<Item = io::Result<Bytes>> + Send
fn decompress_stream<S>(
&self,
compressed_stream: S,
) -> impl Stream<Item = io::Result<Bytes>> + Send
where
S: Stream<Item = io::Result<Bytes>> + Send + 'static + Unpin;
/// Determines whether a file should be compressed based on its MIME type and size
fn should_compress(&self, mime_type: &str, size: u64) -> bool;
}
@@ -77,11 +81,9 @@ pub struct GzipCompressionService {
impl GzipCompressionService {
/// Creates a new service instance
pub fn new() -> Self {
Self {
buffer_pool: None,
}
Self { buffer_pool: None }
}
/// Creates a new service instance with buffer pool
pub fn new_with_buffer_pool(buffer_pool: Arc<BufferPool>) -> Self {
Self {
@@ -98,35 +100,36 @@ impl CompressionService for GzipCompressionService {
if let Some(pool) = &self.buffer_pool {
// Estimate the compression size (approximately 80% of original for typical cases)
let estimated_size = (data.len() as f64 * 0.8) as usize;
// Get a buffer from the pool
let buffer = pool.get_buffer().await;
// Check if the buffer is large enough
if buffer.capacity() >= estimated_size {
// Run compression in a worker thread using the buffer
let buffer_ptr = Arc::new(tokio::sync::Mutex::new(buffer));
let buffer_clone = buffer_ptr.clone();
// Compress data
// Clone the data to avoid lifetime issues
let data_owned = data.to_vec();
let result = tokio::task::spawn_blocking(move || {
let mut encoder = GzEncoderRead::new(&data_owned[..], level.into());
// Try to lock the mutex (should not fail since we are in a separate thread)
let mut buffer_guard = match futures::executor::block_on(buffer_clone.lock()) {
buffer => buffer,
};
// Read directly into the buffer
let read_bytes = encoder.read(buffer_guard.as_mut_slice())?;
buffer_guard.set_used(read_bytes);
Ok(()) as io::Result<()>
}).await;
})
.await;
// Verify result
match result {
Ok(Ok(())) => {
@@ -135,11 +138,11 @@ impl CompressionService for GzipCompressionService {
let cloned_buffer = buffer.clone();
drop(buffer); // Release the mutex first
return Ok(cloned_buffer.into_vec());
},
}
Ok(Err(e)) => {
error!("Compression error with buffer pool: {}", e);
// Fall back to standard implementation
},
}
Err(e) => {
error!("Compression task error with buffer pool: {}", e);
// Fall back to standard implementation
@@ -147,55 +150,58 @@ impl CompressionService for GzipCompressionService {
}
}
}
// Standard implementation if there is no buffer pool or the buffer is insufficient
// Clone the data to avoid lifetime issues
let data_owned = data.to_vec();
tokio::task::spawn_blocking(move || {
let mut encoder = GzEncoderRead::new(&data_owned[..], level.into());
let mut compressed = Vec::new();
encoder.read_to_end(&mut compressed)?;
Ok(compressed)
}).await.unwrap_or_else(|e| {
})
.await
.unwrap_or_else(|e| {
error!("Compression task error: {}", e);
Err(io::Error::other(e.to_string()))
})
}
/// Decompresses data in memory
async fn decompress_data(&self, compressed_data: &[u8]) -> io::Result<Vec<u8>> {
// If we have a buffer pool, use a borrowed buffer for decompression
if let Some(pool) = &self.buffer_pool {
// Estimate the decompression size (approximately 5x of compressed for typical cases)
let estimated_size = compressed_data.len() * 5;
// Get a buffer from the pool
let buffer = pool.get_buffer().await;
// Check if the buffer is large enough
if buffer.capacity() >= estimated_size {
// Clone compressed data to move to the worker
let data = compressed_data.to_vec();
let buffer_ptr = Arc::new(tokio::sync::Mutex::new(buffer));
let buffer_clone = buffer_ptr.clone();
// Decompress data
let result = tokio::task::spawn_blocking(move || {
let mut decoder = GzDecoder::new(&data[..]);
// Try to lock the mutex
let mut buffer_guard = match futures::executor::block_on(buffer_clone.lock()) {
buffer => buffer,
};
// Read directly into the buffer
let read_bytes = decoder.read(buffer_guard.as_mut_slice())?;
buffer_guard.set_used(read_bytes);
Ok(()) as io::Result<()>
}).await;
})
.await;
// Verify result
match result {
Ok(Ok(())) => {
@@ -204,11 +210,11 @@ impl CompressionService for GzipCompressionService {
let cloned_buffer = buffer.clone();
drop(buffer); // Release the mutex first
return Ok(cloned_buffer.into_vec());
},
}
Ok(Err(e)) => {
error!("Decompression error with buffer pool: {}", e);
// Fall back to standard implementation
},
}
Err(e) => {
error!("Decompression task error with buffer pool: {}", e);
// Fall back to standard implementation
@@ -216,7 +222,7 @@ impl CompressionService for GzipCompressionService {
}
}
}
// Standard implementation if there is no buffer pool or the buffer is insufficient
let data = compressed_data.to_vec(); // Clone to move to the worker
tokio::task::spawn_blocking(move || {
@@ -224,26 +230,31 @@ impl CompressionService for GzipCompressionService {
let mut decompressed = Vec::new();
decoder.read_to_end(&mut decompressed)?;
Ok(decompressed)
}).await.unwrap_or_else(|e| {
})
.await
.unwrap_or_else(|e| {
error!("Decompression task error: {}", e);
Err(io::Error::other(e.to_string()))
})
}
/// Compresses a byte stream
fn compress_stream<S>(&self, stream: S, level: CompressionLevel)
-> impl Stream<Item = io::Result<Bytes>> + Send
fn compress_stream<S>(
&self,
stream: S,
level: CompressionLevel,
) -> impl Stream<Item = io::Result<Bytes>> + Send
where
S: Stream<Item = io::Result<Bytes>> + Send + 'static + Unpin
S: Stream<Item = io::Result<Bytes>> + Send + 'static + Unpin,
{
// For now, simplify the implementation to avoid complex pinning issues
// This implementation collects all stream data and then compresses it at once
// Future optimization would be to implement true streaming compression
let compression_level = level;
Box::pin(async_stream::stream! {
let mut data = Vec::new();
// Collect all bytes from the stream
let mut stream = Box::pin(stream);
while let Some(result) = stream.next().await {
@@ -257,7 +268,7 @@ impl CompressionService for GzipCompressionService {
}
}
}
// Compress collected data
match CompressionService::compress_data(self, &data, compression_level).await {
Ok(compressed) => {
@@ -270,19 +281,21 @@ impl CompressionService for GzipCompressionService {
}
})
}
/// Decompresses a byte stream
fn decompress_stream<S>(&self, compressed_stream: S)
-> impl Stream<Item = io::Result<Bytes>> + Send
fn decompress_stream<S>(
&self,
compressed_stream: S,
) -> impl Stream<Item = io::Result<Bytes>> + Send
where
S: Stream<Item = io::Result<Bytes>> + Send + 'static + Unpin
S: Stream<Item = io::Result<Bytes>> + Send + 'static + Unpin,
{
// For now, simplify the implementation to avoid complex pinning issues
// This implementation collects all stream data and then decompresses it at once
// Future optimization would be to implement streaming decompression correctly
Box::pin(async_stream::stream! {
let mut compressed_data = Vec::new();
// Collect all bytes from the stream
let mut stream = Box::pin(compressed_stream);
while let Some(result) = stream.next().await {
@@ -296,7 +309,7 @@ impl CompressionService for GzipCompressionService {
}
}
}
// Decompress collected data
match CompressionService::decompress_data(self, &compressed_data).await {
Ok(decompressed) => {
@@ -309,23 +322,24 @@ impl CompressionService for GzipCompressionService {
}
})
}
/// Determines whether a file should be compressed based on its MIME type and size
fn should_compress(&self, mime_type: &str, size: u64) -> bool {
// Do not compress very small files (overhead)
if size < COMPRESSION_SIZE_THRESHOLD {
return false;
}
// Do not compress already compressed files
if mime_type.starts_with("image/")
&& !mime_type.contains("svg")
&& !mime_type.contains("bmp") {
&& !mime_type.contains("bmp")
{
return false;
}
if mime_type.starts_with("audio/")
|| mime_type.starts_with("video/")
if mime_type.starts_with("audio/")
|| mime_type.starts_with("video/")
|| mime_type.contains("zip")
|| mime_type.contains("gzip")
|| mime_type.contains("compressed")
@@ -341,10 +355,11 @@ impl CompressionService for GzipCompressionService {
|| mime_type.contains("mp3")
|| mime_type.contains("mp4")
|| mime_type.contains("ogg")
|| mime_type.contains("webm") {
|| mime_type.contains("webm")
{
return false;
}
// Compress text files, documents, and other compressible types
true
}
@@ -366,12 +381,20 @@ impl From<PortCompressionLevel> for CompressionLevel {
#[async_trait]
impl CompressionPort for GzipCompressionService {
async fn compress_data(&self, data: &[u8], level: PortCompressionLevel) -> Result<Vec<u8>, DomainError> {
CompressionService::compress_data(self, data, level.into()).await.map_err(DomainError::from)
async fn compress_data(
&self,
data: &[u8],
level: PortCompressionLevel,
) -> Result<Vec<u8>, DomainError> {
CompressionService::compress_data(self, data, level.into())
.await
.map_err(DomainError::from)
}
async fn decompress_data(&self, compressed_data: &[u8]) -> Result<Vec<u8>, DomainError> {
CompressionService::decompress_data(self, compressed_data).await.map_err(DomainError::from)
CompressionService::decompress_data(self, compressed_data)
.await
.map_err(DomainError::from)
}
fn should_compress(&self, mime_type: &str, size: u64) -> bool {
@@ -383,74 +406,111 @@ impl CompressionPort for GzipCompressionService {
mod tests {
use super::*;
use futures::TryStreamExt;
#[tokio::test]
async fn test_compress_decompress_data() {
let service = GzipCompressionService::new();
// Test data
let data = "Hello, world! ".repeat(1000).into_bytes();
// Compress
let compressed = CompressionService::compress_data(&service, &data, CompressionLevel::Default).await.unwrap();
let compressed =
CompressionService::compress_data(&service, &data, CompressionLevel::Default)
.await
.unwrap();
// Verify that compression reduces the size
assert!(compressed.len() < data.len());
// Decompress
let decompressed = CompressionService::decompress_data(&service, &compressed).await.unwrap();
let decompressed = CompressionService::decompress_data(&service, &compressed)
.await
.unwrap();
// Verify that the original data is recovered correctly
assert_eq!(decompressed, data);
}
#[tokio::test]
async fn test_compress_decompress_stream() {
let service = GzipCompressionService::new();
// Create test data
let chunks = vec![
Ok(Bytes::from("Hello, ")),
Ok(Bytes::from("world! ")),
Ok(Bytes::from("This is a test of streaming compression.")),
];
// Convert to stream
let input_stream = futures::stream::iter(chunks);
// Compress the stream
let compressed_stream = service.compress_stream(input_stream, CompressionLevel::Default);
// Collect the compressed bytes
let compressed_bytes = compressed_stream
.try_fold(Vec::new(), |mut acc, chunk| async move {
acc.extend_from_slice(&chunk);
Ok(acc)
}).await.unwrap();
})
.await
.unwrap();
// Decompress the data
let decompressed = CompressionService::decompress_data(&service, &compressed_bytes).await.unwrap();
let decompressed = CompressionService::decompress_data(&service, &compressed_bytes)
.await
.unwrap();
// Verify result
let expected = "Hello, world! This is a test of streaming compression.";
assert_eq!(String::from_utf8(decompressed).unwrap(), expected);
}
#[test]
fn test_should_compress() {
let service = GzipCompressionService::new();
// Cases that should not be compressed
assert!(!CompressionService::should_compress(&service, "image/jpeg", 100 * 1024));
assert!(!CompressionService::should_compress(&service, "video/mp4", 10 * 1024 * 1024));
assert!(!CompressionService::should_compress(&service, "application/zip", 5 * 1024 * 1024));
assert!(!CompressionService::should_compress(
&service,
"image/jpeg",
100 * 1024
));
assert!(!CompressionService::should_compress(
&service,
"video/mp4",
10 * 1024 * 1024
));
assert!(!CompressionService::should_compress(
&service,
"application/zip",
5 * 1024 * 1024
));
// Cases that should be compressed
assert!(CompressionService::should_compress(&service, "text/html", 100 * 1024));
assert!(CompressionService::should_compress(&service, "application/json", 200 * 1024));
assert!(CompressionService::should_compress(&service, "text/plain", 1024 * 1024));
assert!(CompressionService::should_compress(
&service,
"text/html",
100 * 1024
));
assert!(CompressionService::should_compress(
&service,
"application/json",
200 * 1024
));
assert!(CompressionService::should_compress(
&service,
"text/plain",
1024 * 1024
));
// Small files should not be compressed regardless of type
assert!(!CompressionService::should_compress(&service, "text/html", 10 * 1024));
assert!(!CompressionService::should_compress(
&service,
"text/html",
10 * 1024
));
}
}
}
File diff suppressed because it is too large Load Diff
+352 -312
View File
@@ -1,312 +1,352 @@
use bytes::Bytes;
use lru::LruCache;
use std::num::NonZeroUsize;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use tokio::sync::RwLock;
use tracing::{debug, info, warn};
/// Configuration for the file content cache
#[derive(Debug, Clone)]
pub struct FileContentCacheConfig {
/// Maximum size of individual files to cache (bytes)
pub max_file_size: usize,
/// Maximum total cache size (bytes)
pub max_total_size: usize,
/// Maximum number of entries
pub max_entries: usize,
}
impl Default for FileContentCacheConfig {
fn default() -> Self {
Self {
max_file_size: 10 * 1024 * 1024, // 10MB max per file
max_total_size: 512 * 1024 * 1024, // 512MB total cache
max_entries: 10000, // Max 10k files
}
}
}
impl FileContentCacheConfig {
/// Create a new configuration with custom values
pub fn new(max_file_mb: usize, max_total_mb: usize, max_entries: usize) -> Self {
Self {
max_file_size: max_file_mb * 1024 * 1024,
max_total_size: max_total_mb * 1024 * 1024,
max_entries,
}
}
}
/// Cache entry with metadata
#[derive(Clone)]
struct CacheEntry {
content: Bytes,
etag: String,
content_type: String,
}
/// LRU-based file content cache for small/frequently accessed files
///
/// This cache stores the actual content of files in memory for ultra-fast access.
/// It uses an LRU eviction policy and respects memory limits.
pub struct FileContentCache {
cache: RwLock<LruCache<String, CacheEntry>>,
config: FileContentCacheConfig,
current_size: AtomicUsize,
hits: AtomicUsize,
misses: AtomicUsize,
}
impl FileContentCache {
/// Create a new file content cache with the given configuration
pub fn new(config: FileContentCacheConfig) -> Self {
let max_entries = NonZeroUsize::new(config.max_entries).unwrap_or(NonZeroUsize::new(1000).unwrap());
info!(
"Initializing FileContentCache: max_file={}MB, max_total={}MB, max_entries={}",
config.max_file_size / (1024 * 1024),
config.max_total_size / (1024 * 1024),
config.max_entries
);
Self {
cache: RwLock::new(LruCache::new(max_entries)),
config,
current_size: AtomicUsize::new(0),
hits: AtomicUsize::new(0),
misses: AtomicUsize::new(0),
}
}
/// Create a cache with default configuration
pub fn default() -> Self {
Self::new(FileContentCacheConfig::default())
}
/// Check if a file should be cached based on its size
pub fn should_cache(&self, size: usize) -> bool {
size <= self.config.max_file_size
}
/// Get file content from cache
///
/// Returns (content, etag, content_type) if found
pub async fn get(&self, file_id: &str) -> Option<(Bytes, String, String)> {
let mut cache = self.cache.write().await;
if let Some(entry) = cache.get(file_id) {
self.hits.fetch_add(1, Ordering::Relaxed);
debug!("Cache HIT for file: {}", file_id);
return Some((entry.content.clone(), entry.etag.clone(), entry.content_type.clone()));
}
self.misses.fetch_add(1, Ordering::Relaxed);
debug!("Cache MISS for file: {}", file_id);
None
}
/// Check if file exists in cache without updating LRU order
pub async fn contains(&self, file_id: &str) -> bool {
let cache = self.cache.read().await;
cache.contains(file_id)
}
/// Put file content into cache
///
/// Will evict older entries if necessary to make room.
/// Will not cache if file is too large.
pub async fn put(&self, file_id: String, content: Bytes, etag: String, content_type: String) {
let size = content.len();
// Don't cache if too large
if size > self.config.max_file_size {
debug!("File {} too large to cache: {} bytes", file_id, size);
return;
}
// Evict entries until we have room
while self.current_size.load(Ordering::Relaxed) + size > self.config.max_total_size {
let mut cache = self.cache.write().await;
if let Some((evicted_id, evicted_entry)) = cache.pop_lru() {
let evicted_size = evicted_entry.content.len();
self.current_size.fetch_sub(evicted_size, Ordering::Relaxed);
debug!("Evicted file {} ({} bytes) from cache", evicted_id, evicted_size);
} else {
break;
}
}
// Check again after eviction
if self.current_size.load(Ordering::Relaxed) + size > self.config.max_total_size {
warn!("Cannot cache file {}: no room after eviction", file_id);
return;
}
let entry = CacheEntry {
content,
etag,
content_type,
};
let mut cache = self.cache.write().await;
// If replacing an existing entry, subtract its size first
if let Some(old_entry) = cache.peek(&file_id) {
self.current_size.fetch_sub(old_entry.content.len(), Ordering::Relaxed);
}
cache.put(file_id.clone(), entry);
self.current_size.fetch_add(size, Ordering::Relaxed);
debug!("Cached file {} ({} bytes)", file_id, size);
}
/// Remove a file from cache (e.g., when file is deleted or modified)
pub async fn invalidate(&self, file_id: &str) {
let mut cache = self.cache.write().await;
if let Some(entry) = cache.pop(file_id) {
self.current_size.fetch_sub(entry.content.len(), Ordering::Relaxed);
debug!("Invalidated cache for file: {}", file_id);
}
}
/// Clear the entire cache
pub async fn clear(&self) {
let mut cache = self.cache.write().await;
cache.clear();
self.current_size.store(0, Ordering::Relaxed);
info!("Cache cleared");
}
/// Get cache statistics
pub fn stats(&self) -> CacheStats {
let hits = self.hits.load(Ordering::Relaxed);
let misses = self.misses.load(Ordering::Relaxed);
let total = hits + misses;
let hit_rate = if total > 0 {
(hits as f64 / total as f64) * 100.0
} else {
0.0
};
CacheStats {
current_size_bytes: self.current_size.load(Ordering::Relaxed),
max_size_bytes: self.config.max_total_size,
hits,
misses,
hit_rate_percent: hit_rate,
}
}
}
/// Cache statistics
#[derive(Debug, Clone)]
pub struct CacheStats {
pub current_size_bytes: usize,
pub max_size_bytes: usize,
pub hits: usize,
pub misses: usize,
pub hit_rate_percent: f64,
}
/// Thread-safe wrapper for sharing across handlers
pub type SharedFileContentCache = Arc<FileContentCache>;
// ─── ContentCachePort implementation ─────────────────────────
use async_trait::async_trait;
use crate::application::ports::cache_ports::ContentCachePort;
#[async_trait]
impl ContentCachePort for FileContentCache {
fn should_cache(&self, size: usize) -> bool {
FileContentCache::should_cache(self, size)
}
async fn get(&self, file_id: &str) -> Option<(Bytes, String, String)> {
FileContentCache::get(self, file_id).await
}
async fn put(&self, file_id: String, content: Bytes, etag: String, content_type: String) {
FileContentCache::put(self, file_id, content, etag, content_type).await
}
async fn invalidate(&self, file_id: &str) {
FileContentCache::invalidate(self, file_id).await
}
async fn clear(&self) {
FileContentCache::clear(self).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_cache_put_get() {
let cache = FileContentCache::new(FileContentCacheConfig {
max_file_size: 1024,
max_total_size: 4096,
max_entries: 100,
});
let content = Bytes::from("Hello, World!");
cache.put(
"file1".to_string(),
content.clone(),
"etag1".to_string(),
"text/plain".to_string()
).await;
let result = cache.get("file1").await;
assert!(result.is_some());
let (cached_content, etag, content_type) = result.unwrap();
assert_eq!(cached_content, content);
assert_eq!(etag, "etag1");
assert_eq!(content_type, "text/plain");
}
#[tokio::test]
async fn test_cache_eviction() {
let cache = FileContentCache::new(FileContentCacheConfig {
max_file_size: 100,
max_total_size: 200,
max_entries: 100,
});
// Add first file (100 bytes)
let content1 = Bytes::from(vec![0u8; 100]);
cache.put("file1".to_string(), content1, "e1".to_string(), "app/bin".to_string()).await;
// Add second file (100 bytes)
let content2 = Bytes::from(vec![1u8; 100]);
cache.put("file2".to_string(), content2, "e2".to_string(), "app/bin".to_string()).await;
// Add third file - should evict file1
let content3 = Bytes::from(vec![2u8; 100]);
cache.put("file3".to_string(), content3, "e3".to_string(), "app/bin".to_string()).await;
// file1 should be evicted
assert!(cache.get("file1").await.is_none());
// file2 and file3 should exist
assert!(cache.get("file2").await.is_some());
assert!(cache.get("file3").await.is_some());
}
#[tokio::test]
async fn test_cache_invalidate() {
let cache = FileContentCache::new(FileContentCacheConfig::default());
let content = Bytes::from("test");
cache.put("file1".to_string(), content, "e".to_string(), "t".to_string()).await;
assert!(cache.get("file1").await.is_some());
cache.invalidate("file1").await;
assert!(cache.get("file1").await.is_none());
}
}
use bytes::Bytes;
use lru::LruCache;
use std::num::NonZeroUsize;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::RwLock;
use tracing::{debug, info, warn};
/// Configuration for the file content cache
#[derive(Debug, Clone)]
pub struct FileContentCacheConfig {
/// Maximum size of individual files to cache (bytes)
pub max_file_size: usize,
/// Maximum total cache size (bytes)
pub max_total_size: usize,
/// Maximum number of entries
pub max_entries: usize,
}
impl Default for FileContentCacheConfig {
fn default() -> Self {
Self {
max_file_size: 10 * 1024 * 1024, // 10MB max per file
max_total_size: 512 * 1024 * 1024, // 512MB total cache
max_entries: 10000, // Max 10k files
}
}
}
impl FileContentCacheConfig {
/// Create a new configuration with custom values
pub fn new(max_file_mb: usize, max_total_mb: usize, max_entries: usize) -> Self {
Self {
max_file_size: max_file_mb * 1024 * 1024,
max_total_size: max_total_mb * 1024 * 1024,
max_entries,
}
}
}
/// Cache entry with metadata
#[derive(Clone)]
struct CacheEntry {
content: Bytes,
etag: String,
content_type: String,
}
/// LRU-based file content cache for small/frequently accessed files
///
/// This cache stores the actual content of files in memory for ultra-fast access.
/// It uses an LRU eviction policy and respects memory limits.
pub struct FileContentCache {
cache: RwLock<LruCache<String, CacheEntry>>,
config: FileContentCacheConfig,
current_size: AtomicUsize,
hits: AtomicUsize,
misses: AtomicUsize,
}
impl FileContentCache {
/// Create a new file content cache with the given configuration
pub fn new(config: FileContentCacheConfig) -> Self {
let max_entries =
NonZeroUsize::new(config.max_entries).unwrap_or(NonZeroUsize::new(1000).unwrap());
info!(
"Initializing FileContentCache: max_file={}MB, max_total={}MB, max_entries={}",
config.max_file_size / (1024 * 1024),
config.max_total_size / (1024 * 1024),
config.max_entries
);
Self {
cache: RwLock::new(LruCache::new(max_entries)),
config,
current_size: AtomicUsize::new(0),
hits: AtomicUsize::new(0),
misses: AtomicUsize::new(0),
}
}
/// Create a cache with default configuration
pub fn default() -> Self {
Self::new(FileContentCacheConfig::default())
}
/// Check if a file should be cached based on its size
pub fn should_cache(&self, size: usize) -> bool {
size <= self.config.max_file_size
}
/// Get file content from cache
///
/// Returns (content, etag, content_type) if found
pub async fn get(&self, file_id: &str) -> Option<(Bytes, String, String)> {
let mut cache = self.cache.write().await;
if let Some(entry) = cache.get(file_id) {
self.hits.fetch_add(1, Ordering::Relaxed);
debug!("Cache HIT for file: {}", file_id);
return Some((
entry.content.clone(),
entry.etag.clone(),
entry.content_type.clone(),
));
}
self.misses.fetch_add(1, Ordering::Relaxed);
debug!("Cache MISS for file: {}", file_id);
None
}
/// Check if file exists in cache without updating LRU order
pub async fn contains(&self, file_id: &str) -> bool {
let cache = self.cache.read().await;
cache.contains(file_id)
}
/// Put file content into cache
///
/// Will evict older entries if necessary to make room.
/// Will not cache if file is too large.
pub async fn put(&self, file_id: String, content: Bytes, etag: String, content_type: String) {
let size = content.len();
// Don't cache if too large
if size > self.config.max_file_size {
debug!("File {} too large to cache: {} bytes", file_id, size);
return;
}
// Evict entries until we have room
while self.current_size.load(Ordering::Relaxed) + size > self.config.max_total_size {
let mut cache = self.cache.write().await;
if let Some((evicted_id, evicted_entry)) = cache.pop_lru() {
let evicted_size = evicted_entry.content.len();
self.current_size.fetch_sub(evicted_size, Ordering::Relaxed);
debug!(
"Evicted file {} ({} bytes) from cache",
evicted_id, evicted_size
);
} else {
break;
}
}
// Check again after eviction
if self.current_size.load(Ordering::Relaxed) + size > self.config.max_total_size {
warn!("Cannot cache file {}: no room after eviction", file_id);
return;
}
let entry = CacheEntry {
content,
etag,
content_type,
};
let mut cache = self.cache.write().await;
// If replacing an existing entry, subtract its size first
if let Some(old_entry) = cache.peek(&file_id) {
self.current_size
.fetch_sub(old_entry.content.len(), Ordering::Relaxed);
}
cache.put(file_id.clone(), entry);
self.current_size.fetch_add(size, Ordering::Relaxed);
debug!("Cached file {} ({} bytes)", file_id, size);
}
/// Remove a file from cache (e.g., when file is deleted or modified)
pub async fn invalidate(&self, file_id: &str) {
let mut cache = self.cache.write().await;
if let Some(entry) = cache.pop(file_id) {
self.current_size
.fetch_sub(entry.content.len(), Ordering::Relaxed);
debug!("Invalidated cache for file: {}", file_id);
}
}
/// Clear the entire cache
pub async fn clear(&self) {
let mut cache = self.cache.write().await;
cache.clear();
self.current_size.store(0, Ordering::Relaxed);
info!("Cache cleared");
}
/// Get cache statistics
pub fn stats(&self) -> CacheStats {
let hits = self.hits.load(Ordering::Relaxed);
let misses = self.misses.load(Ordering::Relaxed);
let total = hits + misses;
let hit_rate = if total > 0 {
(hits as f64 / total as f64) * 100.0
} else {
0.0
};
CacheStats {
current_size_bytes: self.current_size.load(Ordering::Relaxed),
max_size_bytes: self.config.max_total_size,
hits,
misses,
hit_rate_percent: hit_rate,
}
}
}
/// Cache statistics
#[derive(Debug, Clone)]
pub struct CacheStats {
pub current_size_bytes: usize,
pub max_size_bytes: usize,
pub hits: usize,
pub misses: usize,
pub hit_rate_percent: f64,
}
/// Thread-safe wrapper for sharing across handlers
pub type SharedFileContentCache = Arc<FileContentCache>;
// ─── ContentCachePort implementation ─────────────────────────
use crate::application::ports::cache_ports::ContentCachePort;
use async_trait::async_trait;
#[async_trait]
impl ContentCachePort for FileContentCache {
fn should_cache(&self, size: usize) -> bool {
FileContentCache::should_cache(self, size)
}
async fn get(&self, file_id: &str) -> Option<(Bytes, String, String)> {
FileContentCache::get(self, file_id).await
}
async fn put(&self, file_id: String, content: Bytes, etag: String, content_type: String) {
FileContentCache::put(self, file_id, content, etag, content_type).await
}
async fn invalidate(&self, file_id: &str) {
FileContentCache::invalidate(self, file_id).await
}
async fn clear(&self) {
FileContentCache::clear(self).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_cache_put_get() {
let cache = FileContentCache::new(FileContentCacheConfig {
max_file_size: 1024,
max_total_size: 4096,
max_entries: 100,
});
let content = Bytes::from("Hello, World!");
cache
.put(
"file1".to_string(),
content.clone(),
"etag1".to_string(),
"text/plain".to_string(),
)
.await;
let result = cache.get("file1").await;
assert!(result.is_some());
let (cached_content, etag, content_type) = result.unwrap();
assert_eq!(cached_content, content);
assert_eq!(etag, "etag1");
assert_eq!(content_type, "text/plain");
}
#[tokio::test]
async fn test_cache_eviction() {
let cache = FileContentCache::new(FileContentCacheConfig {
max_file_size: 100,
max_total_size: 200,
max_entries: 100,
});
// Add first file (100 bytes)
let content1 = Bytes::from(vec![0u8; 100]);
cache
.put(
"file1".to_string(),
content1,
"e1".to_string(),
"app/bin".to_string(),
)
.await;
// Add second file (100 bytes)
let content2 = Bytes::from(vec![1u8; 100]);
cache
.put(
"file2".to_string(),
content2,
"e2".to_string(),
"app/bin".to_string(),
)
.await;
// Add third file - should evict file1
let content3 = Bytes::from(vec![2u8; 100]);
cache
.put(
"file3".to_string(),
content3,
"e3".to_string(),
"app/bin".to_string(),
)
.await;
// file1 should be evicted
assert!(cache.get("file1").await.is_none());
// file2 and file3 should exist
assert!(cache.get("file2").await.is_some());
assert!(cache.get("file3").await.is_some());
}
#[tokio::test]
async fn test_cache_invalidate() {
let cache = FileContentCache::new(FileContentCacheConfig::default());
let content = Bytes::from("test");
cache
.put(
"file1".to_string(),
content,
"e".to_string(),
"t".to_string(),
)
.await;
assert!(cache.get("file1").await.is_some());
cache.invalidate("file1").await;
assert!(cache.get("file1").await.is_none());
}
}
+164 -149
View File
@@ -1,3 +1,5 @@
use futures::future::BoxFuture;
use mime_guess::from_path;
use std::collections::{HashMap, VecDeque};
use std::path::{Path, PathBuf};
use std::sync::Arc;
@@ -5,9 +7,7 @@ use std::time::{Duration, Instant, UNIX_EPOCH};
use tokio::fs;
use tokio::sync::RwLock;
use tokio::time;
use futures::future::BoxFuture;
use tracing::debug;
use mime_guess::from_path;
use crate::domain::entities::file::File;
@@ -79,7 +79,7 @@ impl FileMetadata {
ttl: Duration,
) -> Self {
let now = Instant::now();
Self {
path,
exists,
@@ -93,18 +93,18 @@ impl FileMetadata {
access_count: 1,
}
}
/// Updates the last access time
pub fn touch(&mut self) {
self.last_access = Instant::now();
self.access_count += 1;
}
/// Checks if the entry has expired
pub fn is_expired(&self) -> bool {
Instant::now() > self.expires_at
}
/// Updates the expiration time with a new TTL
pub fn update_expiry(&mut self, ttl: Duration) {
self.expires_at = Instant::now() + ttl;
@@ -137,12 +137,12 @@ impl FileMetadataCache {
lru_queue: RwLock::new(VecDeque::with_capacity(max_entries)),
stats: RwLock::new(CacheStats::default()),
config,
ttl_multiplier: 5.0, // Popular entries have 5x TTL
ttl_multiplier: 5.0, // Popular entries have 5x TTL
popularity_threshold: 10, // After 10 accesses it's considered popular
max_entries,
}
}
/// Creates a FileMetadata object from a File object
pub fn create_metadata_from_file(file: &File, abs_path: PathBuf) -> FileMetadata {
let entry_type = CacheEntryType::File;
@@ -150,10 +150,10 @@ impl FileMetadataCache {
let mime_type = Some(file.mime_type().to_string());
let created_at = Some(file.created_at());
let modified_at = Some(file.modified_at());
// Use a standard TTL
let ttl = Duration::from_secs(60); // 1 minute
FileMetadata::new(
abs_path,
true,
@@ -165,147 +165,148 @@ impl FileMetadataCache {
ttl,
)
}
/// Creates a default instance
pub fn default() -> Self {
Self::new(AppConfig::default(), 10_000)
}
/// Creates a cache instance with default configuration
pub fn default_with_config(config: AppConfig) -> Self {
Self::new(config, 50_000) // Larger cache for production system
}
/// Gets file metadata if cached
pub async fn get_metadata(&self, path: &Path) -> Option<FileMetadata> {
let start_time = Instant::now();
let mut cache = self.metadata_cache.write().await;
if let Some(metadata) = cache.get_mut(path) {
// Check if expired
if metadata.is_expired() {
// Remove from cache if expired
cache.remove(path);
// Update statistics
let mut stats = self.stats.write().await;
stats.misses += 1;
stats.expirations += 1;
debug!("Cache entry expired for: {}", path.display());
return None;
}
// Update access time
metadata.touch();
// For popular entries, extend TTL
if metadata.access_count >= self.popularity_threshold {
let new_ttl = match metadata.entry_type {
CacheEntryType::File => Duration::from_millis(
(self.config.timeouts.file_operation_ms as f64 * self.ttl_multiplier) as u64
(self.config.timeouts.file_operation_ms as f64 * self.ttl_multiplier)
as u64,
),
CacheEntryType::Directory => Duration::from_millis(
(self.config.timeouts.dir_operation_ms as f64 * self.ttl_multiplier) as u64
(self.config.timeouts.dir_operation_ms as f64 * self.ttl_multiplier) as u64,
),
_ => Duration::from_secs(60), // 1 minute by default
};
metadata.update_expiry(new_ttl);
debug!("Extended TTL for popular entry: {}", path.display());
}
// Calculate approximate time saved
let elapsed = start_time.elapsed().as_millis() as u64;
let estimated_io_time: u64 = 10; // We assume 10ms minimum for IO operation
let time_saved = estimated_io_time.saturating_sub(elapsed);
// Update statistics
let mut stats = self.stats.write().await;
stats.hits += 1;
stats.time_saved_ms += time_saved;
debug!("Cache hit for: {}", path.display());
// Also keep the LRU queue updated
self.update_lru(path.to_path_buf()).await;
// Clone to return
return Some(metadata.clone());
}
// Not found in cache
let mut stats = self.stats.write().await;
stats.misses += 1;
debug!("Cache miss for: {}", path.display());
None
}
/// Updates the LRU queue
async fn update_lru(&self, path: PathBuf) {
let mut lru = self.lru_queue.write().await;
// Remove if already exists
if let Some(pos) = lru.iter().position(|p| p == &path) {
lru.remove(pos);
}
// Add to the end (most recent)
lru.push_back(path);
}
/// Checks if a file exists
pub async fn exists(&self, path: &Path) -> Option<bool> {
if let Some(metadata) = self.get_metadata(path).await {
return Some(metadata.exists);
}
None
}
/// Checks if a path is a directory
pub async fn is_dir(&self, path: &Path) -> Option<bool> {
if let Some(metadata) = self.get_metadata(path).await {
return Some(metadata.entry_type == CacheEntryType::Directory);
}
None
}
/// Checks if a path is a file
pub async fn is_file(&self, path: &Path) -> Option<bool> {
if let Some(metadata) = self.get_metadata(path).await {
return Some(metadata.entry_type == CacheEntryType::File);
}
None
}
/// Gets the size of a file
pub async fn get_size(&self, path: &Path) -> Option<u64> {
if let Some(metadata) = self.get_metadata(path).await {
return metadata.size;
}
None
}
/// Gets the MIME type of a file
pub async fn get_mime_type(&self, path: &Path) -> Option<String> {
if let Some(metadata) = self.get_metadata(path).await {
return metadata.mime_type;
}
None
}
/// Refreshes metadata for a path
pub async fn refresh_metadata(&self, path: &Path) -> Result<FileMetadata, std::io::Error> {
// Perform actual filesystem read
let metadata = fs::metadata(path).await?;
// Determine entry type
let entry_type = if metadata.is_dir() {
CacheEntryType::Directory
@@ -314,37 +315,47 @@ impl FileMetadataCache {
} else {
CacheEntryType::Unknown
};
// Get size for files
let size = if metadata.is_file() {
Some(metadata.len())
} else {
None
};
// Get MIME type for files
let mime_type = if metadata.is_file() {
Some(from_path(path).first_or_octet_stream().to_string())
} else {
None
};
// Get timestamps
let created_at = metadata.created()
.map(|time| time.duration_since(UNIX_EPOCH).unwrap_or_default().as_secs())
let created_at = metadata
.created()
.map(|time| {
time.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
})
.ok();
let modified_at = metadata.modified()
.map(|time| time.duration_since(UNIX_EPOCH).unwrap_or_default().as_secs())
let modified_at = metadata
.modified()
.map(|time| {
time.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
})
.ok();
// Determine appropriate TTL
let ttl = if metadata.is_dir() {
Duration::from_millis(self.config.timeouts.dir_operation_ms)
} else {
Duration::from_millis(self.config.timeouts.file_operation_ms)
};
// Create metadata entry
let file_metadata = FileMetadata::new(
path.to_path_buf(),
@@ -356,50 +367,50 @@ impl FileMetadataCache {
modified_at,
ttl,
);
// Update cache
self.update_cache(file_metadata.clone()).await;
Ok(file_metadata)
}
/// Updates the cache with new metadata
pub async fn update_cache(&self, metadata: FileMetadata) {
// Avoid full cache before inserting
self.ensure_capacity().await;
let path = metadata.path.clone();
// Insert into cache
{
let mut cache = self.metadata_cache.write().await;
cache.insert(path.clone(), metadata);
// Update statistics
let mut stats = self.stats.write().await;
stats.inserts += 1;
}
// Update the LRU queue
self.update_lru(path).await;
}
/// Ensures there is space in the cache
async fn ensure_capacity(&self) {
let cache_size = {
let cache = self.metadata_cache.read().await;
cache.len()
};
if cache_size >= self.max_entries {
self.evict_lru_entries(cache_size / 10).await; // Free up 10%
}
}
/// Removes least recently used entries
async fn evict_lru_entries(&self, count: usize) {
let mut paths_to_remove = Vec::with_capacity(count);
// Get entries to remove from the LRU queue
{
let mut lru = self.lru_queue.write().await;
@@ -411,7 +422,7 @@ impl FileMetadataCache {
}
}
}
// Remove from the main cache
{
let mut cache = self.metadata_cache.write().await;
@@ -419,22 +430,22 @@ impl FileMetadataCache {
cache.remove(&path);
}
}
debug!("Evicted {} LRU entries from cache", count);
}
/// Invalidate a specific cache entry
pub async fn invalidate(&self, path: &Path) {
// Remove from the main cache
{
let mut cache = self.metadata_cache.write().await;
cache.remove(path);
// Update statistics
let mut stats = self.stats.write().await;
stats.invalidations += 1;
}
// Remove from the LRU queue
let path_buf = path.to_path_buf();
{
@@ -443,15 +454,15 @@ impl FileMetadataCache {
lru.remove(pos);
}
}
debug!("Invalidated cache entry for: {}", path.display());
}
/// Recursively invalidate entries under a directory
pub async fn invalidate_directory(&self, dir_path: &Path) {
let dir_str = dir_path.to_string_lossy().to_string();
let mut paths_to_remove = Vec::new();
// Find all paths that start with the directory
{
let cache = self.metadata_cache.read().await;
@@ -462,32 +473,32 @@ impl FileMetadataCache {
}
}
}
// Update statistics
{
let mut stats = self.stats.write().await;
stats.invalidations += paths_to_remove.len();
}
// Remove each found path
for path in paths_to_remove {
self.invalidate(&path).await;
}
debug!("Invalidated directory and contents: {}", dir_path.display());
}
/// Get current cache statistics
pub async fn get_stats(&self) -> CacheStats {
let stats = self.stats.read().await;
stats.clone()
}
/// Clears all expired entries from the cache
pub async fn clear_expired(&self) {
let now = Instant::now();
let mut paths_to_remove = Vec::new();
// Find expired entries
{
let cache = self.metadata_cache.read().await;
@@ -497,112 +508,116 @@ impl FileMetadataCache {
}
}
}
// Update statistics
{
let mut stats = self.stats.write().await;
stats.expirations += paths_to_remove.len();
}
// Save the number of entries for logging
let num_paths = paths_to_remove.len();
// Remove expired entries
for path in paths_to_remove {
self.invalidate(&path).await;
}
debug!("Cleared {} expired entries from cache", num_paths);
}
/// Starts the periodic cleanup process
pub fn start_cleanup_task(cache: Arc<Self>) -> BoxFuture<'static, ()> {
Box::pin(async move {
let cleanup_interval = Duration::from_secs(60); // Every minute
loop {
// Wait for the interval
time::sleep(cleanup_interval).await;
// Clean expired entries
cache.clear_expired().await;
// Log statistics
let stats = cache.get_stats().await;
let cache_size = {
let cache_map = cache.metadata_cache.read().await;
cache_map.len()
};
debug!(
"Cache stats: size={}, hits={}, misses={}, hit_ratio={:.2}%, time_saved={}ms",
cache_size,
stats.hits,
stats.misses,
if stats.hits + stats.misses > 0 {
if stats.hits + stats.misses > 0 {
(stats.hits as f64 * 100.0) / (stats.hits + stats.misses) as f64
} else {
0.0
} else {
0.0
},
stats.time_saved_ms
);
}
})
}
/// Preloads metadata for entire directories (useful for initialization)
pub async fn preload_directory(&self, dir_path: &Path, recursive: bool, max_depth: usize) -> Result<usize, std::io::Error> {
self._preload_directory_internal(dir_path, recursive, max_depth, 0).await
pub async fn preload_directory(
&self,
dir_path: &Path,
recursive: bool,
max_depth: usize,
) -> Result<usize, std::io::Error> {
self._preload_directory_internal(dir_path, recursive, max_depth, 0)
.await
}
/// Internal preload implementation with depth tracking
async fn _preload_directory_internal(
&self,
dir_path: &Path,
recursive: bool,
max_depth: usize,
current_depth: usize
&self,
dir_path: &Path,
recursive: bool,
max_depth: usize,
current_depth: usize,
) -> Result<usize, std::io::Error> {
Box::pin(async move {
if current_depth > max_depth {
return Ok(0);
}
// Get directory entries
let mut entries = fs::read_dir(dir_path).await?;
let mut count = 0;
// Process each entry
while let Some(entry) = entries.next_entry().await? {
let path = entry.path();
let metadata = fs::metadata(&path).await?;
// Refresh metadata for this entry
self.refresh_metadata(&path).await?;
count += 1;
// Recursively process subdirectories if needed
if recursive && metadata.is_dir() {
// Box to break recursion
count += self._preload_directory_internal(
&path,
recursive,
max_depth,
current_depth + 1
).await?;
if current_depth > max_depth {
return Ok(0);
}
}
Ok(count)
}).await
// Get directory entries
let mut entries = fs::read_dir(dir_path).await?;
let mut count = 0;
// Process each entry
while let Some(entry) = entries.next_entry().await? {
let path = entry.path();
let metadata = fs::metadata(&path).await?;
// Refresh metadata for this entry
self.refresh_metadata(&path).await?;
count += 1;
// Recursively process subdirectories if needed
if recursive && metadata.is_dir() {
// Box to break recursion
count += self
._preload_directory_internal(&path, recursive, max_depth, current_depth + 1)
.await?;
}
}
Ok(count)
})
.await
}
}
// ─── MetadataCachePort implementation ────────────────────────
use async_trait::async_trait;
use crate::application::ports::cache_ports::{MetadataCachePort, CachedMetadataDto};
use crate::application::ports::cache_ports::{CachedMetadataDto, MetadataCachePort};
use crate::common::errors::DomainError;
use async_trait::async_trait;
#[async_trait]
impl MetadataCachePort for FileMetadataCache {
@@ -654,46 +669,46 @@ mod tests {
use tempfile::tempdir;
use tokio::fs::File;
use tokio::io::AsyncWriteExt;
#[tokio::test]
async fn test_cache_operations() {
// Create temporary directory for tests
let temp_dir = tempdir().unwrap();
let file_path = temp_dir.path().join("test_file.txt");
// Create a test file
let mut file = File::create(&file_path).await.unwrap();
file.write_all(b"test content").await.unwrap();
file.flush().await.unwrap();
drop(file);
// Create cache
let config = AppConfig::default();
let cache = FileMetadataCache::new(config, 1000);
// Verify initial miss
assert!(cache.exists(&file_path).await.is_none());
// Refresh and verify hit
let metadata = cache.refresh_metadata(&file_path).await.unwrap();
assert_eq!(metadata.entry_type, CacheEntryType::File);
assert_eq!(metadata.size, Some(12)); // "test content" = 12 bytes
// Verify it now exists in cache
assert_eq!(cache.exists(&file_path).await, Some(true));
assert_eq!(cache.is_file(&file_path).await, Some(true));
// Invalidate and verify it no longer exists in cache
cache.invalidate(&file_path).await;
assert!(cache.exists(&file_path).await.is_none());
// Verify statistics
let stats = cache.get_stats().await;
assert_eq!(stats.inserts, 1);
assert_eq!(stats.invalidations, 1);
assert!(stats.hits > 0);
}
#[tokio::test]
async fn test_directory_operations() {
// Create directory structure for tests
@@ -702,33 +717,33 @@ mod tests {
let base_path = temp_dir.path().canonicalize().unwrap();
let sub_dir = base_path.join("subdir");
fs::create_dir(&sub_dir).await.unwrap();
let file1 = base_path.join("file1.txt");
let file2 = sub_dir.join("file2.txt");
File::create(&file1).await.unwrap();
File::create(&file2).await.unwrap();
// Create cache
let config = AppConfig::default();
let cache = FileMetadataCache::new(config, 1000);
// Preload directory recursively
// preload_directory caches the *contents* of the directory, not the root itself
let count = cache.preload_directory(&base_path, true, 2).await.unwrap();
assert_eq!(count, 3); // subdir, file1, file2
// Verify existence in cache (only contents, not the root)
assert_eq!(cache.is_dir(&sub_dir).await, Some(true));
assert_eq!(cache.is_file(&file1).await, Some(true));
assert_eq!(cache.is_file(&file2).await, Some(true));
// Invalidate directory and contents
cache.invalidate_directory(&base_path).await;
// Verify nothing exists in cache
assert!(cache.exists(&sub_dir).await.is_none());
assert!(cache.exists(&file1).await.is_none());
assert!(cache.exists(&file2).await.is_none());
}
}
}
@@ -1,17 +1,17 @@
use async_trait::async_trait;
use serde_json::Value;
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::RwLock;
use async_trait::async_trait;
use serde_json::Value;
use tokio::fs;
use crate::domain::services::i18n_service::{I18nService, I18nError, I18nResult, Locale};
use crate::domain::services::i18n_service::{I18nError, I18nResult, I18nService, Locale};
/// File system implementation of the I18nService
pub struct FileSystemI18nService {
/// Base directory containing translation files
translations_dir: PathBuf,
/// Cached translations (locale code -> JSON data)
cache: RwLock<HashMap<Locale, Value>>,
}
@@ -32,17 +32,18 @@ impl FileSystemI18nService {
cache: RwLock::new(HashMap::new()),
}
}
/// Get translation file path for a locale
fn get_locale_file_path(&self, locale: Locale) -> PathBuf {
self.translations_dir.join(format!("{}.json", locale.as_str()))
self.translations_dir
.join(format!("{}.json", locale.as_str()))
}
/// Get a nested key from JSON data
fn get_nested_value(&self, data: &Value, key: &str) -> Option<String> {
let parts: Vec<&str> = key.split('.').collect();
let mut current = data;
for part in &parts[0..parts.len() - 1] {
if let Some(next) = current.get(part) {
current = next;
@@ -50,13 +51,14 @@ impl FileSystemI18nService {
return None;
}
}
if let Some(last_part) = parts.last()
&& let Some(value) = current.get(last_part)
&& value.is_string() {
return value.as_str().map(|s| s.to_string());
}
&& value.is_string()
{
return value.as_str().map(|s| s.to_string());
}
None
}
}
@@ -71,73 +73,86 @@ impl I18nService for FileSystemI18nService {
if let Some(value) = self.get_nested_value(translations, key) {
return Ok(value);
}
// Try to use English as fallback if we couldn't find the key
if locale != Locale::English
&& let Some(english_translations) = cache.get(&Locale::English)
&& let Some(value) = self.get_nested_value(english_translations, key) {
return Ok(value);
}
&& let Some(value) = self.get_nested_value(english_translations, key)
{
return Ok(value);
}
return Err(I18nError::KeyNotFound(key.to_string()));
}
}
// If not cached, load translations and try again
self.load_translations(locale).await?;
{
let cache = self.cache.read().unwrap();
if let Some(translations) = cache.get(&locale) {
if let Some(value) = self.get_nested_value(translations, key) {
return Ok(value);
}
// Try to use English as fallback
if locale != Locale::English
&& let Some(english_translations) = cache.get(&Locale::English)
&& let Some(value) = self.get_nested_value(english_translations, key) {
return Ok(value);
}
&& let Some(value) = self.get_nested_value(english_translations, key)
{
return Ok(value);
}
}
}
Err(I18nError::KeyNotFound(key.to_string()))
}
async fn load_translations(&self, locale: Locale) -> I18nResult<()> {
let file_path = self.get_locale_file_path(locale);
tracing::info!("Loading translations for locale {} from {:?}", locale.as_str(), file_path);
tracing::info!(
"Loading translations for locale {} from {:?}",
locale.as_str(),
file_path
);
// Check if file exists
if !file_path.exists() {
return Err(I18nError::InvalidLocale(locale.as_str().to_string()));
}
// Read and parse file
let content = fs::read_to_string(&file_path)
.await
.map_err(|e| I18nError::LoadError(format!("Failed to read translation file: {}", e)))?;
let translations: Value = serde_json::from_str(&content)
.map_err(|e| I18nError::LoadError(format!("Failed to parse translation file: {}", e)))?;
let translations: Value = serde_json::from_str(&content).map_err(|e| {
I18nError::LoadError(format!("Failed to parse translation file: {}", e))
})?;
// Update cache
{
let mut cache = self.cache.write().unwrap();
cache.insert(locale, translations);
}
tracing::info!("Translations loaded for locale {}", locale.as_str());
Ok(())
}
async fn available_locales(&self) -> Vec<Locale> {
vec![Locale::English, Locale::Spanish, Locale::French, Locale::German, Locale::Portuguese]
vec![
Locale::English,
Locale::Spanish,
Locale::French,
Locale::German,
Locale::Portuguese,
]
}
async fn is_supported(&self, locale: Locale) -> bool {
let file_path = self.get_locale_file_path(locale);
file_path.exists()
}
}
}
+140 -94
View File
@@ -1,9 +1,9 @@
use tokio::fs::{self, OpenOptions, File};
use tokio::io::AsyncWriteExt;
use std::path::Path;
use std::io::Error as IoError;
use std::path::Path;
use tempfile::NamedTempFile;
use tracing::{warn, error};
use tokio::fs::{self, File, OpenOptions};
use tokio::io::AsyncWriteExt;
use tracing::{error, warn};
/// Utility functions for file system operations with proper synchronization
pub struct FileSystemUtils;
@@ -13,59 +13,73 @@ impl FileSystemUtils {
/// Uses a safe atomic write pattern: write to temp file, fsync, rename
pub async fn atomic_write<P: AsRef<Path>>(path: P, contents: &[u8]) -> Result<(), IoError> {
let path = path.as_ref();
// Ensure parent directory exists
if let Some(parent) = path.parent() {
fs::create_dir_all(parent).await?;
}
// Create a temporary file in the same directory
let dir = path.parent().unwrap_or_else(|| Path::new("."));
let temp_file = match NamedTempFile::new_in(dir) {
Ok(file) => file,
Err(e) => {
error!("Failed to create temporary file in {}: {}", dir.display(), e);
return Err(IoError::other(format!("Failed to create temporary file: {}", e)));
error!(
"Failed to create temporary file in {}: {}",
dir.display(),
e
);
return Err(IoError::other(format!(
"Failed to create temporary file: {}",
e
)));
}
};
let temp_path = temp_file.path().to_path_buf();
// Convert to tokio file and write contents
let std_file = temp_file.as_file().try_clone()?;
let mut file = File::from_std(std_file);
file.write_all(contents).await?;
// Ensure data is synced to disk
file.flush().await?;
file.sync_all().await?;
// Rename the temporary file to the target path (atomic operation on most filesystems)
fs::rename(&temp_path, path).await?;
// Sync the directory to ensure the rename is persisted
if let Some(parent) = path.parent() {
match Self::sync_directory(parent).await {
Ok(_) => {},
Ok(_) => {}
Err(e) => {
warn!("Failed to sync directory {}: {}. File was written but directory entry might not be durable.",
parent.display(), e);
warn!(
"Failed to sync directory {}: {}. File was written but directory entry might not be durable.",
parent.display(),
e
);
}
}
}
Ok(())
}
/// Creates or appends to a file with fsync
pub async fn write_with_sync<P: AsRef<Path>>(path: P, contents: &[u8], append: bool) -> Result<(), IoError> {
pub async fn write_with_sync<P: AsRef<Path>>(
path: P,
contents: &[u8],
append: bool,
) -> Result<(), IoError> {
let path = path.as_ref();
// Ensure parent directory exists
if let Some(parent) = path.parent() {
fs::create_dir_all(parent).await?;
}
// Open file with appropriate options
let mut file = OpenOptions::new()
.write(true)
@@ -74,140 +88,162 @@ impl FileSystemUtils {
.append(append)
.open(path)
.await?;
// Write contents
file.write_all(contents).await?;
// Ensure data is synced to disk
file.flush().await?;
file.sync_all().await?;
Ok(())
}
/// Creates directories with fsync
pub async fn create_dir_with_sync<P: AsRef<Path>>(path: P) -> Result<(), IoError> {
let path = path.as_ref();
// Create directory
fs::create_dir_all(path).await?;
// Sync the directory
Self::sync_directory(path).await?;
// Sync parent directory to ensure directory creation is persisted
if let Some(parent) = path.parent() {
match Self::sync_directory(parent).await {
Ok(_) => {},
Ok(_) => {}
Err(e) => {
warn!("Failed to sync parent directory {}: {}. Directory was created but entry might not be durable.",
parent.display(), e);
warn!(
"Failed to sync parent directory {}: {}. Directory was created but entry might not be durable.",
parent.display(),
e
);
}
}
}
Ok(())
}
/// Renames a file or directory with proper syncing
pub async fn rename_with_sync<P: AsRef<Path>, Q: AsRef<Path>>(from: P, to: Q) -> Result<(), IoError> {
pub async fn rename_with_sync<P: AsRef<Path>, Q: AsRef<Path>>(
from: P,
to: Q,
) -> Result<(), IoError> {
let from = from.as_ref();
let to = to.as_ref();
// Ensure parent directory of destination exists
if let Some(parent) = to.parent() {
fs::create_dir_all(parent).await?;
}
// Perform rename
fs::rename(from, to).await?;
// Sync parent directories to ensure rename is persisted
if let Some(from_parent) = from.parent() {
match Self::sync_directory(from_parent).await {
Ok(_) => {},
Ok(_) => {}
Err(e) => {
warn!("Failed to sync source directory {}: {}. Rename completed but might not be durable.",
from_parent.display(), e);
warn!(
"Failed to sync source directory {}: {}. Rename completed but might not be durable.",
from_parent.display(),
e
);
}
}
}
if let Some(to_parent) = to.parent() {
match Self::sync_directory(to_parent).await {
Ok(_) => {},
Ok(_) => {}
Err(e) => {
warn!("Failed to sync destination directory {}: {}. Rename completed but might not be durable.",
to_parent.display(), e);
warn!(
"Failed to sync destination directory {}: {}. Rename completed but might not be durable.",
to_parent.display(),
e
);
}
}
}
Ok(())
}
/// Removes a file with directory syncing
pub async fn remove_file_with_sync<P: AsRef<Path>>(path: P) -> Result<(), IoError> {
let path = path.as_ref();
// Remove file
fs::remove_file(path).await?;
// Sync parent directory to ensure removal is persisted
if let Some(parent) = path.parent() {
match Self::sync_directory(parent).await {
Ok(_) => {},
Ok(_) => {}
Err(e) => {
warn!("Failed to sync directory after file removal {}: {}. File was removed but entry might not be durable.",
parent.display(), e);
warn!(
"Failed to sync directory after file removal {}: {}. File was removed but entry might not be durable.",
parent.display(),
e
);
}
}
}
Ok(())
}
/// Removes a directory with parent directory syncing
pub async fn remove_dir_with_sync<P: AsRef<Path>>(path: P, recursive: bool) -> Result<(), IoError> {
pub async fn remove_dir_with_sync<P: AsRef<Path>>(
path: P,
recursive: bool,
) -> Result<(), IoError> {
let path = path.as_ref();
// Remove directory
if recursive {
fs::remove_dir_all(path).await?;
} else {
fs::remove_dir(path).await?;
}
// Sync parent directory to ensure removal is persisted
if let Some(parent) = path.parent() {
match Self::sync_directory(parent).await {
Ok(_) => {},
Ok(_) => {}
Err(e) => {
warn!("Failed to sync directory after directory removal {}: {}. Directory was removed but entry might not be durable.",
parent.display(), e);
warn!(
"Failed to sync directory after directory removal {}: {}. Directory was removed but entry might not be durable.",
parent.display(),
e
);
}
}
}
Ok(())
}
/// Syncs a directory to ensure its contents are durable
async fn sync_directory<P: AsRef<Path>>(path: P) -> Result<(), IoError> {
let path = path.as_ref();
// Open directory with read permissions
let dir_file = match OpenOptions::new()
.read(true)
.open(path)
.await {
Ok(file) => file,
Err(e) => {
warn!("Failed to open directory for syncing {}: {}", path.display(), e);
return Err(e);
}
};
let dir_file = match OpenOptions::new().read(true).open(path).await {
Ok(file) => file,
Err(e) => {
warn!(
"Failed to open directory for syncing {}: {}",
path.display(),
e
);
return Err(e);
}
};
// Sync the directory
dir_file.sync_all().await
}
@@ -219,62 +255,72 @@ mod tests {
use tempfile::tempdir;
use tokio::fs;
use tokio::io::AsyncReadExt;
#[tokio::test]
async fn test_atomic_write() {
let temp_dir = tempdir().unwrap();
let file_path = temp_dir.path().join("test.txt");
// Write data atomically
FileSystemUtils::atomic_write(&file_path, b"Hello, world!").await.unwrap();
FileSystemUtils::atomic_write(&file_path, b"Hello, world!")
.await
.unwrap();
// Read back the data
let mut file = fs::File::open(&file_path).await.unwrap();
let mut contents = String::new();
file.read_to_string(&mut contents).await.unwrap();
assert_eq!(contents, "Hello, world!");
}
#[tokio::test]
async fn test_write_with_sync() {
let temp_dir = tempdir().unwrap();
let file_path = temp_dir.path().join("test.txt");
// Write data with sync
FileSystemUtils::write_with_sync(&file_path, b"First line\n", false).await.unwrap();
FileSystemUtils::write_with_sync(&file_path, b"First line\n", false)
.await
.unwrap();
// Append data
FileSystemUtils::write_with_sync(&file_path, b"Second line", true).await.unwrap();
FileSystemUtils::write_with_sync(&file_path, b"Second line", true)
.await
.unwrap();
// Read back the data
let mut file = fs::File::open(&file_path).await.unwrap();
let mut contents = String::new();
file.read_to_string(&mut contents).await.unwrap();
assert_eq!(contents, "First line\nSecond line");
}
#[tokio::test]
async fn test_rename_with_sync() {
let temp_dir = tempdir().unwrap();
let source_path = temp_dir.path().join("source.txt");
let dest_path = temp_dir.path().join("dest.txt");
// Create source file
FileSystemUtils::write_with_sync(&source_path, b"Test content", false).await.unwrap();
FileSystemUtils::write_with_sync(&source_path, b"Test content", false)
.await
.unwrap();
// Rename file
FileSystemUtils::rename_with_sync(&source_path, &dest_path).await.unwrap();
FileSystemUtils::rename_with_sync(&source_path, &dest_path)
.await
.unwrap();
// Verify source doesn't exist
assert!(!source_path.exists());
// Verify destination exists
let mut file = fs::File::open(&dest_path).await.unwrap();
let mut contents = String::new();
file.read_to_string(&mut contents).await.unwrap();
assert_eq!(contents, "Test content");
}
}
}
@@ -1,14 +1,14 @@
use async_trait::async_trait;
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{Mutex, RwLock, Semaphore};
use tracing::{debug, error, info, warn};
use async_trait::async_trait;
use crate::domain::services::path_service::StoragePath;
use crate::infrastructure::services::id_mapping_service::{IdMappingService, IdMappingError};
use crate::common::errors::DomainError;
use crate::application::ports::outbound::IdMappingPort;
use crate::common::errors::DomainError;
use crate::domain::services::path_service::StoragePath;
use crate::infrastructure::services::id_mapping_service::{IdMappingError, IdMappingService};
/// Maximum number of entries in the cache
const MAX_CACHE_SIZE: usize = 10_000;
@@ -20,19 +20,19 @@ const CACHE_TTL_SECONDS: u64 = 60 * 5; // 5 minutes
pub struct IdMappingOptimizer {
/// Base ID mapping service
base_service: Arc<IdMappingService>,
/// Path to ID cache (path -> id)
path_to_id_cache: RwLock<HashMap<String, (String, Instant)>>,
/// ID to path cache (id -> path)
id_to_path_cache: RwLock<HashMap<String, (String, Instant)>>,
/// Hit counter
stats: RwLock<OptimizerStats>,
/// Semaphore to limit batch operations
batch_limiter: Semaphore,
/// Pending batch queue
pending_batch: Mutex<BatchQueue>,
}
@@ -44,17 +44,17 @@ pub struct OptimizerStats {
pub path_by_id_queries: usize,
/// Number of cache hits for get_path_by_id
pub path_by_id_hits: usize,
/// Total number of get_or_create_id queries
pub get_id_queries: usize,
/// Number of cache hits for get_or_create_id
pub get_id_hits: usize,
/// Number of batch operations performed
pub batch_operations: usize,
/// Total number of IDs processed in batch
pub batch_items_processed: usize,
/// Last cache cleanup timestamp
pub last_cleanup: Option<Instant>,
}
@@ -68,7 +68,6 @@ struct BatchQueue {
id_to_path_requests: HashSet<String>,
}
/// Result of a batch operation
struct BatchResult {
/// Path to ID mapping
@@ -89,85 +88,93 @@ impl IdMappingOptimizer {
pending_batch: Mutex::new(BatchQueue::default()),
}
}
/// Gets optimizer statistics
pub async fn get_stats(&self) -> OptimizerStats {
self.stats.read().await.clone()
}
/// Cleans expired cache entries
pub async fn cleanup_cache(&self) {
let now = Instant::now();
let ttl = Duration::from_secs(CACHE_TTL_SECONDS);
// Clean path_to_id cache
{
let mut cache = self.path_to_id_cache.write().await;
let initial_size = cache.len();
// Retain only non-expired entries
cache.retain(|_, (_, timestamp)| {
now.duration_since(*timestamp) < ttl
});
cache.retain(|_, (_, timestamp)| now.duration_since(*timestamp) < ttl);
let removed = initial_size - cache.len();
if removed > 0 {
debug!("Cleaned {} expired entries from path_to_id cache", removed);
}
}
// Clean id_to_path cache
{
let mut cache = self.id_to_path_cache.write().await;
let initial_size = cache.len();
// Retain only non-expired entries
cache.retain(|_, (_, timestamp)| {
now.duration_since(*timestamp) < ttl
});
cache.retain(|_, (_, timestamp)| now.duration_since(*timestamp) < ttl);
let removed = initial_size - cache.len();
if removed > 0 {
debug!("Cleaned {} expired entries from id_to_path cache", removed);
}
}
// Update statistics
{
let mut stats = self.stats.write().await;
stats.last_cleanup = Some(now);
}
}
/// Starts periodic cleanup task
pub fn start_cleanup_task(optimizer: Arc<Self>) {
tokio::spawn(async move {
let cleanup_interval = Duration::from_secs(CACHE_TTL_SECONDS / 2);
loop {
tokio::time::sleep(cleanup_interval).await;
optimizer.cleanup_cache().await;
// Log statistics periodically
let stats = optimizer.get_stats().await;
info!("ID Mapping Optimizer stats - Path queries: {}, hits: {} ({}%), ID queries: {}, hits: {} ({}%), Batch ops: {}, items: {}",
info!(
"ID Mapping Optimizer stats - Path queries: {}, hits: {} ({}%), ID queries: {}, hits: {} ({}%), Batch ops: {}, items: {}",
stats.path_by_id_queries,
stats.path_by_id_hits,
if stats.path_by_id_queries > 0 { stats.path_by_id_hits as f64 * 100.0 / stats.path_by_id_queries as f64 } else { 0.0 },
if stats.path_by_id_queries > 0 {
stats.path_by_id_hits as f64 * 100.0 / stats.path_by_id_queries as f64
} else {
0.0
},
stats.get_id_queries,
stats.get_id_hits,
if stats.get_id_queries > 0 { stats.get_id_hits as f64 * 100.0 / stats.get_id_queries as f64 } else { 0.0 },
if stats.get_id_queries > 0 {
stats.get_id_hits as f64 * 100.0 / stats.get_id_queries as f64
} else {
0.0
},
stats.batch_operations,
stats.batch_items_processed
);
}
});
}
/// Adds a request to the pending queue for batch processing
async fn queue_path_to_id_request(&self, path: &StoragePath) -> Result<Option<String>, IdMappingError> {
async fn queue_path_to_id_request(
&self,
path: &StoragePath,
) -> Result<Option<String>, IdMappingError> {
let path_str = path.to_string();
// Check first in the cache
{
let cache = self.path_to_id_cache.read().await;
@@ -177,42 +184,42 @@ impl IdMappingOptimizer {
let mut stats = self.stats.write().await;
stats.get_id_hits += 1;
}
return Ok(Some(id.clone()));
}
}
// If not in cache, add to batch queue
{
let mut batch_queue = self.pending_batch.lock().await;
batch_queue.path_to_id_requests.insert(path_str);
}
// Not found in cache, must be processed in batch
Ok(None)
}
/// Processes pending requests in batch
async fn process_batch(&self) -> Result<BatchResult, IdMappingError> {
// Acquire permit for batch operation
let _permit = self.batch_limiter.acquire().await.unwrap();
// Get pending requests
let (path_requests, id_requests) = {
let mut batch_queue = self.pending_batch.lock().await;
let paths = std::mem::take(&mut batch_queue.path_to_id_requests);
let ids = std::mem::take(&mut batch_queue.id_to_path_requests);
(paths, ids)
};
// Create results
let mut result = BatchResult {
path_to_id: HashMap::with_capacity(path_requests.len()),
id_to_path: HashMap::with_capacity(id_requests.len()),
};
// Process path->id requests in batch
for path_str in path_requests {
let path = StoragePath::from_string(&path_str);
@@ -220,14 +227,14 @@ impl IdMappingOptimizer {
Ok(id) => {
result.path_to_id.insert(path_str.clone(), id.clone());
result.id_to_path.insert(id, path_str);
},
}
Err(e) => {
error!("Error batch-processing path {}: {}", path_str, e);
// Continue with remaining requests
}
}
}
// Process id->path requests in batch
for id in id_requests {
match self.base_service.get_path_by_id(&id).await {
@@ -235,37 +242,37 @@ impl IdMappingOptimizer {
let path_str = path.to_string();
result.id_to_path.insert(id.clone(), path_str.clone());
result.path_to_id.insert(path_str, id);
},
}
Err(e) => {
error!("Error batch-processing ID {}: {}", id, e);
// Continue with remaining requests
}
}
}
// Update cache with batch results
{
let mut path_cache = self.path_to_id_cache.write().await;
let mut id_cache = self.id_to_path_cache.write().await;
let now = Instant::now();
for (path, id) in &result.path_to_id {
path_cache.insert(path.clone(), (id.clone(), now));
}
for (id, path) in &result.id_to_path {
id_cache.insert(id.clone(), (path.clone(), now));
}
}
// Update statistics
{
let mut stats = self.stats.write().await;
stats.batch_operations += 1;
stats.batch_items_processed += result.path_to_id.len() + result.id_to_path.len();
}
// Save changes to disk in the background
let service_clone = self.base_service.clone();
tokio::spawn(async move {
@@ -273,36 +280,37 @@ impl IdMappingOptimizer {
error!("Error saving ID mapping changes: {}", e);
}
});
Ok(result)
}
/// Forces processing of pending requests if there are enough
async fn trigger_batch_if_needed(&self, min_batch_size: usize) -> Result<(), IdMappingError> {
// Check if there are enough pending requests
let should_process = {
let batch_queue = self.pending_batch.lock().await;
batch_queue.path_to_id_requests.len() + batch_queue.id_to_path_requests.len() >= min_batch_size
batch_queue.path_to_id_requests.len() + batch_queue.id_to_path_requests.len()
>= min_batch_size
};
// Process if necessary
if should_process {
self.process_batch().await?;
}
Ok(())
}
/// Preload a set of paths to get their IDs in batch
pub async fn preload_paths(&self, paths: Vec<StoragePath>) -> Result<(), IdMappingError> {
// Only proceed if there are paths to load
if paths.is_empty() {
return Ok(());
}
// Paths we need to load (those not in cache)
let mut paths_to_load = Vec::new();
// Check cache first
{
let cache = self.path_to_id_cache.read().await;
@@ -313,12 +321,12 @@ impl IdMappingOptimizer {
}
}
}
// If all were in cache, finish
if paths_to_load.is_empty() {
return Ok(());
}
// Add paths to queue for batch processing
{
let mut batch_queue = self.pending_batch.lock().await;
@@ -326,23 +334,23 @@ impl IdMappingOptimizer {
batch_queue.path_to_id_requests.insert(path);
}
}
// Execute batch processing immediately
self.process_batch().await?;
Ok(())
}
/// Preload a set of IDs to get their paths in batch
pub async fn preload_ids(&self, ids: Vec<String>) -> Result<(), IdMappingError> {
// Only proceed if there are IDs to load
if ids.is_empty() {
return Ok(());
}
// IDs we need to load (those not in cache)
let mut ids_to_load = Vec::new();
// Check cache first
{
let cache = self.id_to_path_cache.read().await;
@@ -352,12 +360,12 @@ impl IdMappingOptimizer {
}
}
}
// If all were in cache, finish
if ids_to_load.is_empty() {
return Ok(());
}
// Add IDs to queue for batch processing
{
let mut batch_queue = self.pending_batch.lock().await;
@@ -365,10 +373,10 @@ impl IdMappingOptimizer {
batch_queue.id_to_path_requests.insert(id);
}
}
// Execute batch processing immediately
self.process_batch().await?;
Ok(())
}
}
@@ -381,9 +389,9 @@ impl IdMappingPort for IdMappingOptimizer {
let mut stats = self.stats.write().await;
stats.get_id_queries += 1;
}
let path_str = path.to_string();
// Check cache first
{
let cache = self.path_to_id_cache.read().await;
@@ -393,55 +401,61 @@ impl IdMappingPort for IdMappingOptimizer {
let mut stats = self.stats.write().await;
stats.get_id_hits += 1;
}
return Ok(id.clone());
}
}
// If not in cache, try adding to batch queue first
let queued_result = self.queue_path_to_id_request(path).await?;
if let Some(id) = queued_result {
return Ok(id);
}
// Trigger batch processing if enough items accumulated
self.trigger_batch_if_needed(20).await?;
// Try to get from the base service
let id = self.base_service.get_or_create_id(path).await?;
// Update cache with the new ID
{
let mut path_cache = self.path_to_id_cache.write().await;
let mut id_cache = self.id_to_path_cache.write().await;
let now = Instant::now();
// Control cache size
if path_cache.len() >= MAX_CACHE_SIZE {
warn!("Path-to-ID cache size reached limit ({}), clearing oldest entries", MAX_CACHE_SIZE);
warn!(
"Path-to-ID cache size reached limit ({}), clearing oldest entries",
MAX_CACHE_SIZE
);
path_cache.clear();
}
if id_cache.len() >= MAX_CACHE_SIZE {
warn!("ID-to-path cache size reached limit ({}), clearing oldest entries", MAX_CACHE_SIZE);
warn!(
"ID-to-path cache size reached limit ({}), clearing oldest entries",
MAX_CACHE_SIZE
);
id_cache.clear();
}
path_cache.insert(path_str.clone(), (id.clone(), now));
id_cache.insert(id.clone(), (path_str, now));
}
Ok(id)
}
async fn get_path_by_id(&self, id: &str) -> Result<StoragePath, DomainError> {
// Update statistics
{
let mut stats = self.stats.write().await;
stats.path_by_id_queries += 1;
}
// Check first in the cache
{
let cache = self.id_to_path_cache.read().await;
@@ -451,92 +465,98 @@ impl IdMappingPort for IdMappingOptimizer {
let mut stats = self.stats.write().await;
stats.path_by_id_hits += 1;
}
return Ok(StoragePath::from_string(path_str));
}
}
// Get from the base service
let path = self.base_service.get_path_by_id(id).await?;
// Update cache
{
let mut id_cache = self.id_to_path_cache.write().await;
let mut path_cache = self.path_to_id_cache.write().await;
let now = Instant::now();
let path_str = path.to_string();
// Control cache size
if id_cache.len() >= MAX_CACHE_SIZE {
warn!("ID-to-path cache size reached limit ({}), clearing oldest entries", MAX_CACHE_SIZE);
warn!(
"ID-to-path cache size reached limit ({}), clearing oldest entries",
MAX_CACHE_SIZE
);
id_cache.clear();
}
if path_cache.len() >= MAX_CACHE_SIZE {
warn!("Path-to-ID cache size reached limit ({}), clearing oldest entries", MAX_CACHE_SIZE);
warn!(
"Path-to-ID cache size reached limit ({}), clearing oldest entries",
MAX_CACHE_SIZE
);
path_cache.clear();
}
id_cache.insert(id.to_string(), (path_str.clone(), now));
path_cache.insert(path_str, (id.to_string(), now));
}
Ok(path)
}
async fn update_path(&self, id: &str, new_path: &StoragePath) -> Result<(), DomainError> {
// Invalidate cache for this ID
{
let mut id_cache = self.id_to_path_cache.write().await;
let mut path_cache = self.path_to_id_cache.write().await;
// Remove the ID entry
if let Some((old_path, _)) = id_cache.remove(id) {
path_cache.remove(&old_path);
}
}
// Update in the base service
let result = self.base_service.update_path(id, new_path).await?;
// Update cache with new mapping
{
let mut id_cache = self.id_to_path_cache.write().await;
let mut path_cache = self.path_to_id_cache.write().await;
let now = Instant::now();
let path_str = new_path.to_string();
id_cache.insert(id.to_string(), (path_str.clone(), now));
path_cache.insert(path_str, (id.to_string(), now));
}
Ok(result)
}
async fn remove_id(&self, id: &str) -> Result<(), DomainError> {
// Invalidate cache for this ID
{
let mut id_cache = self.id_to_path_cache.write().await;
let mut path_cache = self.path_to_id_cache.write().await;
// Remove the ID entry
if let Some((path, _)) = id_cache.remove(id) {
path_cache.remove(&path);
}
}
// Remove from the base service
self.base_service.remove_id(id).await?;
Ok(())
}
async fn save_changes(&self) -> Result<(), DomainError> {
// Delegate to the base service
self.base_service.save_changes().await?;
Ok(())
}
}
@@ -545,90 +565,99 @@ impl IdMappingPort for IdMappingOptimizer {
mod tests {
use super::*;
use tempfile::tempdir;
async fn create_test_service() -> (Arc<IdMappingService>, Arc<IdMappingOptimizer>) {
let temp_dir = tempdir().unwrap();
let map_path = temp_dir.path().join("id_map.json");
let base_service = Arc::new(IdMappingService::new(map_path).await.unwrap());
let optimizer = Arc::new(IdMappingOptimizer::new(base_service.clone()));
(base_service, optimizer)
}
#[tokio::test]
async fn test_basic_caching() {
let (_, optimizer) = create_test_service().await;
let path = StoragePath::from_string("/test/file.txt");
// First call should use the base service
let id = optimizer.get_or_create_id(&path).await.unwrap();
assert!(!id.is_empty(), "ID should not be empty");
// Second call should use cache
let id2 = optimizer.get_or_create_id(&path).await.unwrap();
assert_eq!(id, id2, "Same path should return same ID");
// Verify cache statistics
let stats = optimizer.get_stats().await;
assert_eq!(stats.get_id_queries, 2, "Should have 2 queries");
assert_eq!(stats.get_id_hits, 1, "Should have 1 hit");
}
#[tokio::test]
async fn test_batch_processing() {
let (_, optimizer) = create_test_service().await;
// Create a batch of paths
let mut paths = Vec::new();
for i in 0..50 {
paths.push(StoragePath::from_string(&format!("/test/batch/file{}.txt", i)));
paths.push(StoragePath::from_string(&format!(
"/test/batch/file{}.txt",
i
)));
}
// Preload the paths
optimizer.preload_paths(paths.clone()).await.unwrap();
// Verify all are in cache
for path in &paths {
let id = optimizer.get_or_create_id(path).await.unwrap();
assert!(!id.is_empty(), "ID should be available for path");
}
// Verify statistics
let stats = optimizer.get_stats().await;
assert_eq!(stats.batch_operations, 1, "Should have 1 batch operation");
assert!(stats.batch_items_processed >= 50, "Should have processed at least 50 items");
assert!(
stats.batch_items_processed >= 50,
"Should have processed at least 50 items"
);
// Verify all subsequent queries are cache hits
assert_eq!(stats.get_id_hits, 50, "All subsequente queries should be cache hits");
assert_eq!(
stats.get_id_hits, 50,
"All subsequente queries should be cache hits"
);
}
#[tokio::test]
async fn test_cache_cleanup() {
let (_, optimizer) = create_test_service().await;
// Create some entries
let path = StoragePath::from_string("/test/cleanup.txt");
let id = optimizer.get_or_create_id(&path).await.unwrap();
// Verify initial statistics
{
let stats = optimizer.get_stats().await;
assert_eq!(stats.get_id_queries, 1, "Should have 1 query");
assert_eq!(stats.get_id_hits, 0, "Should have 0 hits");
}
// Run cleanup (should not remove anything yet)
optimizer.cleanup_cache().await;
// Verify cache is still working
let id2 = optimizer.get_or_create_id(&path).await.unwrap();
assert_eq!(id, id2, "Cache should still work after cleanup");
{
let stats = optimizer.get_stats().await;
assert_eq!(stats.get_id_hits, 1, "Should have 1 hit after cleanup");
}
}
}
}
+296 -196
View File
@@ -1,32 +1,32 @@
use std::path::PathBuf;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use tokio::sync::{RwLock, Mutex};
use std::path::PathBuf;
use tokio::fs;
use tokio::sync::{Mutex, RwLock};
use tokio::time;
use uuid::Uuid;
use serde::{Serialize, Deserialize};
use async_trait::async_trait;
use crate::domain::services::path_service::StoragePath;
use crate::common::errors::{DomainError, ErrorKind};
use crate::application::ports::outbound::IdMappingPort;
use crate::common::config::TimeoutConfig;
use crate::common::errors::{DomainError, ErrorKind};
use crate::domain::services::path_service::StoragePath;
/// Specific error for the ID mapping service
#[derive(Debug, thiserror::Error)]
pub enum IdMappingError {
#[error("ID not found: {0}")]
NotFound(String),
#[error("IO error: {0}")]
IoError(#[from] std::io::Error),
#[error("Timeout error: {0}")]
Timeout(String),
#[error("Serialization error: {0}")]
SerializationError(#[from] serde_json::Error),
#[error("Other error: {0}")]
Other(String),
}
@@ -39,21 +39,22 @@ impl From<IdMappingError> for DomainError {
IdMappingError::IoError(e) => DomainError::new(
ErrorKind::InternalError,
"IdMapping",
format!("IO error: {}", e)
).with_source(e),
IdMappingError::Timeout(msg) => DomainError::timeout(
"IdMapping",
format!("Timeout: {}", msg)
),
format!("IO error: {}", e),
)
.with_source(e),
IdMappingError::Timeout(msg) => {
DomainError::timeout("IdMapping", format!("Timeout: {}", msg))
}
IdMappingError::SerializationError(e) => DomainError::new(
ErrorKind::InternalError,
"IdMapping",
format!("Serialization error: {}", e)
).with_source(e),
format!("Serialization error: {}", e),
)
.with_source(e),
IdMappingError::Other(msg) => DomainError::new(
ErrorKind::InternalError,
"IdMapping",
format!("Other error: {}", msg)
format!("Other error: {}", msg),
),
}
}
@@ -64,7 +65,7 @@ impl From<IdMappingError> for DomainError {
struct IdMap {
path_to_id: HashMap<String, String>,
id_to_path: HashMap<String, String>, // Field for efficient bidirectional lookup
version: u32, // Version to detect changes
version: u32, // Version to detect changes
}
/// Service to manage mappings between paths and unique IDs
@@ -81,7 +82,7 @@ impl IdMappingService {
pub async fn new(map_path: PathBuf) -> Result<Self, DomainError> {
let timeouts = TimeoutConfig::default();
let id_map = Self::load_id_map(&map_path, &timeouts).await?;
Ok(Self {
map_path,
id_map: RwLock::new(id_map),
@@ -90,7 +91,7 @@ impl IdMappingService {
pending_save: RwLock::new(false),
})
}
/// Creates an in-memory ID mapping service (for testing)
///
/// Similar functionality as new_in_memory but with a simpler signature for dummy use
@@ -103,7 +104,7 @@ impl IdMappingService {
pending_save: RwLock::new(false),
}
}
/// Creates an in-memory ID mapping service (for testing - original version)
pub fn new_in_memory() -> Self {
Self {
@@ -114,19 +115,30 @@ impl IdMappingService {
pending_save: RwLock::new(false),
}
}
/// Loads the ID map from disk with robust error handling
async fn load_id_map(map_path: &PathBuf, timeouts: &TimeoutConfig) -> Result<IdMap, DomainError> {
async fn load_id_map(
map_path: &PathBuf,
timeouts: &TimeoutConfig,
) -> Result<IdMap, DomainError> {
if map_path.exists() {
// Try to read with timeout to avoid indefinite blocking
let read_result = time::timeout(
timeouts.lock_timeout(),
fs::read_to_string(map_path)
).await
.map_err(|_| DomainError::timeout("IdMapping", format!("Timeout reading ID map from {}", map_path.display())))?;
let content = read_result.map_err(|e| DomainError::internal_error("IdMapping", format!("Failed to read ID map from {}: {}", map_path.display(), e)))?;
let read_result = time::timeout(timeouts.lock_timeout(), fs::read_to_string(map_path))
.await
.map_err(|_| {
DomainError::timeout(
"IdMapping",
format!("Timeout reading ID map from {}", map_path.display()),
)
})?;
let content = read_result.map_err(|e| {
DomainError::internal_error(
"IdMapping",
format!("Failed to read ID map from {}: {}", map_path.display(), e),
)
})?;
// Parse the JSON
match serde_json::from_str::<IdMap>(&content) {
Ok(mut map) => {
@@ -139,11 +151,14 @@ impl IdMappingService {
}
tracing::info!("Rebuilt inverse mapping with {} entries", rebuild_count);
}
tracing::info!("Loaded ID map with {} entries (version: {})",
map.path_to_id.len(), map.version);
tracing::info!(
"Loaded ID map with {} entries (version: {})",
map.path_to_id.len(),
map.version
);
return Ok(map);
},
}
Err(e) => {
tracing::error!("Error parsing ID map: {}", e);
// Try to backup the corrupted file
@@ -153,7 +168,7 @@ impl IdMappingService {
} else {
tracing::info!("Backed up corrupted ID map to {}", backup_path.display());
}
tracing::info!("Creating new empty map after error");
return Ok(IdMap {
path_to_id: HashMap::new(),
@@ -163,7 +178,7 @@ impl IdMappingService {
}
}
}
// Return an empty map if the file doesn't exist and create the file
tracing::info!("No existing ID map found, creating new empty map");
let empty_map = IdMap {
@@ -171,217 +186,259 @@ impl IdMappingService {
id_to_path: HashMap::new(),
version: 1, // Start with version 1
};
// Ensure directory exists
if let Some(parent) = map_path.parent()
&& !parent.exists()
&& let Err(e) = fs::create_dir_all(parent).await {
tracing::error!("Failed to create directory for ID map: {}", e);
}
&& let Err(e) = fs::create_dir_all(parent).await
{
tracing::error!("Failed to create directory for ID map: {}", e);
}
// Write empty map to file (best-effort: the in-memory map is valid even if disk write fails)
match serde_json::to_string_pretty(&empty_map) {
Ok(json) => {
if let Err(e) = fs::write(map_path, json).await {
tracing::warn!("Could not write initial empty ID map (will retry on next save): {}", e);
tracing::warn!(
"Could not write initial empty ID map (will retry on next save): {}",
e
);
} else {
tracing::info!("Created initial empty ID map at {}", map_path.display());
}
},
}
Err(e) => {
tracing::error!("Failed to serialize empty ID map: {}", e);
}
}
Ok(empty_map)
}
/// Saves the ID map to disk safely
async fn save_id_map(&self) -> Result<(), DomainError> {
// Acquire exclusive lock for saving
let _lock = time::timeout(
self.timeouts.lock_timeout(),
self.save_mutex.lock()
).await
.map_err(|_| DomainError::timeout("IdMapping", "Timeout acquiring save lock for ID mapping"))?;
let _lock = time::timeout(self.timeouts.lock_timeout(), self.save_mutex.lock())
.await
.map_err(|_| {
DomainError::timeout("IdMapping", "Timeout acquiring save lock for ID mapping")
})?;
// Create JSON with read lock to minimize lock hold time
let json = {
let mut map = time::timeout(
self.timeouts.lock_timeout(),
self.id_map.write()
).await
.map_err(|_| DomainError::timeout("IdMapping", "Timeout acquiring write lock for ID mapping"))?;
let mut map = time::timeout(self.timeouts.lock_timeout(), self.id_map.write())
.await
.map_err(|_| {
DomainError::timeout("IdMapping", "Timeout acquiring write lock for ID mapping")
})?;
// Increment version only if there are pending changes to save
let pending = *self.pending_save.read().await;
if pending {
map.version += 1;
tracing::debug!("Incrementing ID map version to {}", map.version);
}
// Use serde with reasonably safe defaults
serde_json::to_string_pretty(&*map)
.map_err(|e| DomainError::internal_error("IdMapping", format!("Failed to serialize ID map to JSON: {}", e)))?
serde_json::to_string_pretty(&*map).map_err(|e| {
DomainError::internal_error(
"IdMapping",
format!("Failed to serialize ID map to JSON: {}", e),
)
})?
};
// Write to a temporary file first to avoid corruption
let temp_path = self.map_path.with_extension("json.tmp");
fs::write(&temp_path, &json).await
.map_err(|e| DomainError::internal_error("IdMapping", format!("Failed to write temporary ID map to {}: {}", temp_path.display(), e)))?;
fs::write(&temp_path, &json).await.map_err(|e| {
DomainError::internal_error(
"IdMapping",
format!(
"Failed to write temporary ID map to {}: {}",
temp_path.display(),
e
),
)
})?;
// Perform the atomic rename
fs::rename(&temp_path, &self.map_path).await
.map_err(|e| DomainError::internal_error("IdMapping", format!("Failed to rename temporary ID map to {}: {}", self.map_path.display(), e)))?;
fs::rename(&temp_path, &self.map_path).await.map_err(|e| {
DomainError::internal_error(
"IdMapping",
format!(
"Failed to rename temporary ID map to {}: {}",
self.map_path.display(),
e
),
)
})?;
// Reset pending flag
{
let mut pending = self.pending_save.write().await;
*pending = false;
}
tracing::info!("Saved ID map successfully to {}", self.map_path.display());
Ok(())
}
/// Generates a unique ID
fn generate_id(&self) -> String {
Uuid::new_v4().to_string()
}
/// Marks changes as pending
async fn mark_pending(&self) {
let mut pending = self.pending_save.write().await;
*pending = true;
}
/// Gets the ID for a path or generates a new one if it doesn't exist
pub async fn get_or_create_id(&self, path: &StoragePath) -> Result<String, IdMappingError> {
let path_str = path.to_string();
// First attempt with read lock (more efficient)
{
let read_result = match time::timeout(
self.timeouts.lock_timeout(),
self.id_map.read()
).await {
Ok(guard) => guard,
Err(_) => return Err(IdMappingError::Timeout("Timeout acquiring read lock for ID mapping".to_string())),
};
let read_result =
match time::timeout(self.timeouts.lock_timeout(), self.id_map.read()).await {
Ok(guard) => guard,
Err(_) => {
return Err(IdMappingError::Timeout(
"Timeout acquiring read lock for ID mapping".to_string(),
));
}
};
if let Some(id) = read_result.path_to_id.get(&path_str) {
return Ok(id.clone());
}
}
// If not found, acquire write lock
let write_result = match time::timeout(
self.timeouts.lock_timeout(),
self.id_map.write()
).await {
Ok(guard) => guard,
Err(_) => return Err(IdMappingError::Timeout("Timeout acquiring write lock for ID mapping".to_string())),
};
let write_result =
match time::timeout(self.timeouts.lock_timeout(), self.id_map.write()).await {
Ok(guard) => guard,
Err(_) => {
return Err(IdMappingError::Timeout(
"Timeout acquiring write lock for ID mapping".to_string(),
));
}
};
let mut map = write_result;
// Check again (it could have been added while we were waiting for the lock)
if let Some(id) = map.path_to_id.get(&path_str) {
return Ok(id.clone());
}
// Generate a new ID and store it
let id = self.generate_id();
map.path_to_id.insert(path_str.clone(), id.clone());
map.id_to_path.insert(id.clone(), path_str);
// Mark as pending for saving
drop(map); // Release the write lock before acquiring another
self.mark_pending().await;
tracing::debug!("Created new ID mapping: {} -> {}", path.to_string(), id);
Ok(id)
}
/// Gets a path by its ID with timeout handling
pub async fn get_path_by_id(&self, id: &str) -> Result<StoragePath, IdMappingError> {
let read_result = match time::timeout(
self.timeouts.lock_timeout(),
self.id_map.read()
).await {
Ok(guard) => guard,
Err(_) => return Err(IdMappingError::Timeout("Timeout acquiring read lock for ID lookup".to_string())),
};
let read_result =
match time::timeout(self.timeouts.lock_timeout(), self.id_map.read()).await {
Ok(guard) => guard,
Err(_) => {
return Err(IdMappingError::Timeout(
"Timeout acquiring read lock for ID lookup".to_string(),
));
}
};
if let Some(path_str) = read_result.id_to_path.get(id) {
return Ok(StoragePath::from_string(path_str));
}
Err(IdMappingError::NotFound(id.to_string()))
}
/// Updates the mapping of an existing ID to a new path
pub async fn update_path(&self, id: &str, new_path: &StoragePath) -> Result<(), IdMappingError> {
let write_result = match time::timeout(
self.timeouts.lock_timeout(),
self.id_map.write()
).await {
Ok(guard) => guard,
Err(_) => return Err(IdMappingError::Timeout("Timeout acquiring write lock for ID update".to_string())),
};
pub async fn update_path(
&self,
id: &str,
new_path: &StoragePath,
) -> Result<(), IdMappingError> {
let write_result =
match time::timeout(self.timeouts.lock_timeout(), self.id_map.write()).await {
Ok(guard) => guard,
Err(_) => {
return Err(IdMappingError::Timeout(
"Timeout acquiring write lock for ID update".to_string(),
));
}
};
let mut map = write_result;
// Find the previous path to remove it
if let Some(old_path) = map.id_to_path.get(id).cloned() {
map.path_to_id.remove(&old_path);
// Register the new path
let new_path_str = new_path.to_string();
map.path_to_id.insert(new_path_str.clone(), id.to_string());
map.id_to_path.insert(id.to_string(), new_path_str);
// Mark as pending
drop(map); // Release the write lock before acquiring another
self.mark_pending().await;
tracing::debug!("Updated path mapping for ID {}: {} -> {}",
id, old_path, new_path.to_string());
tracing::debug!(
"Updated path mapping for ID {}: {} -> {}",
id,
old_path,
new_path.to_string()
);
Ok(())
} else {
Err(IdMappingError::NotFound(id.to_string()))
}
}
/// Removes an ID from the map
pub async fn remove_id(&self, id: &str) -> Result<(), IdMappingError> {
let write_result = match time::timeout(
self.timeouts.lock_timeout(),
self.id_map.write()
).await {
Ok(guard) => guard,
Err(_) => return Err(IdMappingError::Timeout("Timeout acquiring write lock for ID removal".to_string())),
};
let write_result =
match time::timeout(self.timeouts.lock_timeout(), self.id_map.write()).await {
Ok(guard) => guard,
Err(_) => {
return Err(IdMappingError::Timeout(
"Timeout acquiring write lock for ID removal".to_string(),
));
}
};
let mut map = write_result;
// Find the path to remove it
if let Some(path) = map.id_to_path.remove(id) {
map.path_to_id.remove(&path);
// Mark as pending
drop(map); // Release the write lock before acquiring another
self.mark_pending().await;
tracing::debug!("Removed ID mapping: {} -> {}", id, path);
Ok(())
} else {
Err(IdMappingError::NotFound(id.to_string()))
}
}
/// Saves pending changes to disk immediately, without debounce
pub async fn save_pending_changes(&self) -> Result<(), IdMappingError> {
// Check if there are pending changes
@@ -391,50 +448,64 @@ impl IdMappingService {
return Ok(());
}
}
// Save immediately (without debounce or spawn)
match self.save_id_map().await {
Ok(_) => {
tracing::info!("ID mappings saved successfully to disk at {}", self.map_path.display());
tracing::info!(
"ID mappings saved successfully to disk at {}",
self.map_path.display()
);
// Explicitly verify that the file exists and has size
match std::fs::metadata(&self.map_path) {
Ok(metadata) => {
if metadata.len() > 0 {
tracing::info!("Verified saved map file exists with size: {} bytes", metadata.len());
tracing::info!(
"Verified saved map file exists with size: {} bytes",
metadata.len()
);
} else {
tracing::warn!("Map file exists but has zero size - this might cause issues");
tracing::warn!(
"Map file exists but has zero size - this might cause issues"
);
}
},
}
Err(e) => {
tracing::error!("Failed to verify saved map file: {}", e);
// Try a second save if verification fails
if let Err(retry_err) = self.save_id_map().await {
tracing::error!("Second save attempt also failed: {}", retry_err);
return Err(IdMappingError::IoError(std::io::Error::other(
format!("Failed to verify and retry save: {}", retry_err)
)));
return Err(IdMappingError::IoError(std::io::Error::other(format!(
"Failed to verify and retry save: {}",
retry_err
))));
}
tracing::info!("Second save attempt succeeded");
}
}
Ok(())
},
}
Err(e) => {
tracing::error!("Failed to save ID map to {}: {}", self.map_path.display(), e);
tracing::error!(
"Failed to save ID map to {}: {}",
self.map_path.display(),
e
);
// Try a second save with delay in case of error
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
match self.save_id_map().await {
Ok(_) => {
tracing::info!("Second save attempt succeeded after initial failure");
Ok(())
},
}
Err(retry_e) => {
tracing::error!("Second save attempt also failed: {}", retry_e);
Err(IdMappingError::IoError(std::io::Error::other(
format!("Failed to save ID mappings after retry: {}", retry_e)
)))
Err(IdMappingError::IoError(std::io::Error::other(format!(
"Failed to save ID mappings after retry: {}",
retry_e
))))
}
}
}
@@ -446,32 +517,58 @@ impl IdMappingService {
impl IdMappingPort for IdMappingService {
/// Gets the ID for a path or generates a new one if it doesn't exist
async fn get_or_create_id(&self, path: &StoragePath) -> Result<String, DomainError> {
self.get_or_create_id(path).await
.map_err(|e| DomainError::internal_error("IdMapping", format!("Failed to get or create ID for path: {}: {}", path.to_string(), e)))
self.get_or_create_id(path).await.map_err(|e| {
DomainError::internal_error(
"IdMapping",
format!(
"Failed to get or create ID for path: {}: {}",
path.to_string(),
e
),
)
})
}
/// Gets a path by its ID with timeout handling
async fn get_path_by_id(&self, id: &str) -> Result<StoragePath, DomainError> {
self.get_path_by_id(id).await
.map_err(|e| DomainError::internal_error("IdMapping", format!("Failed to get path for ID: {}: {}", id, e)))
self.get_path_by_id(id).await.map_err(|e| {
DomainError::internal_error(
"IdMapping",
format!("Failed to get path for ID: {}: {}", id, e),
)
})
}
/// Updates the mapping of an existing ID to a new path
async fn update_path(&self, id: &str, new_path: &StoragePath) -> Result<(), DomainError> {
self.update_path(id, new_path).await
.map_err(|e| DomainError::internal_error("IdMapping", format!("Failed to update path for ID: {} to {}: {}", id, new_path.to_string(), e)))
self.update_path(id, new_path).await.map_err(|e| {
DomainError::internal_error(
"IdMapping",
format!(
"Failed to update path for ID: {} to {}: {}",
id,
new_path.to_string(),
e
),
)
})
}
/// Removes an ID from the map
async fn remove_id(&self, id: &str) -> Result<(), DomainError> {
self.remove_id(id).await
.map_err(|e| DomainError::internal_error("IdMapping", format!("Failed to remove ID: {}: {}", id, e)))
self.remove_id(id).await.map_err(|e| {
DomainError::internal_error("IdMapping", format!("Failed to remove ID: {}: {}", id, e))
})
}
/// Saves pending changes to disk
async fn save_changes(&self) -> Result<(), DomainError> {
self.save_pending_changes().await
.map_err(|e| DomainError::internal_error("IdMapping", format!("Failed to save pending ID mapping changes: {}", e)))
self.save_pending_changes().await.map_err(|e| {
DomainError::internal_error(
"IdMapping",
format!("Failed to save pending ID mapping changes: {}", e),
)
})
}
}
@@ -501,7 +598,7 @@ impl Clone for IdMappingService {
Self {
map_path: self.map_path.clone(),
id_map: RwLock::new(IdMap::default()), // This is not used in the async task
save_mutex: Mutex::new(()), // Neither is this
save_mutex: Mutex::new(()), // Neither is this
timeouts: self.timeouts.clone(),
pending_save: RwLock::new(false),
}
@@ -513,100 +610,103 @@ mod tests {
use super::*;
use std::time::Duration;
use tempfile::tempdir;
#[tokio::test]
async fn test_get_or_create_id() {
let temp_dir = tempdir().unwrap();
let map_path = temp_dir.path().join("id_map.json");
let service = IdMappingService::new(map_path).await.unwrap();
let path = StoragePath::from_string("/test/file.txt");
let id = service.get_or_create_id(&path).await.unwrap();
assert!(!id.is_empty(), "ID should not be empty");
// Verify that the same ID is returned for the same path
let id2 = service.get_or_create_id(&path).await.unwrap();
assert_eq!(id, id2, "Same path should return same ID");
}
#[tokio::test]
async fn test_update_path() {
let temp_dir = tempdir().unwrap();
let map_path = temp_dir.path().join("id_map.json");
let service = IdMappingService::new(map_path).await.unwrap();
let old_path = StoragePath::from_string("/test/old.txt");
let id = service.get_or_create_id(&old_path).await.unwrap();
let new_path = StoragePath::from_string("/test/new.txt");
service.update_path(&id, &new_path).await.unwrap();
let retrieved_path = service.get_path_by_id(&id).await.unwrap();
assert_eq!(retrieved_path, new_path, "Path should be updated");
}
#[tokio::test]
async fn test_save_and_load() {
let temp_dir = tempdir().unwrap();
let map_path = temp_dir.path().join("id_map.json");
// Create and populate the service
let service = IdMappingService::new(map_path.clone()).await.unwrap();
let path1 = StoragePath::from_string("/test/file1.txt");
let path2 = StoragePath::from_string("/test/file2.txt");
let id1 = service.get_or_create_id(&path1).await.unwrap();
let id2 = service.get_or_create_id(&path2).await.unwrap();
// Save changes
service.save_pending_changes().await.unwrap();
// Wait to ensure the async save completes
tokio::time::sleep(Duration::from_millis(500)).await;
// Create a new service that should load the same map
let service2 = IdMappingService::new(map_path).await.unwrap();
// Verify that the IDs match
let loaded_id1 = service2.get_or_create_id(&path1).await.unwrap();
let loaded_id2 = service2.get_or_create_id(&path2).await.unwrap();
assert_eq!(id1, loaded_id1, "ID1 should be preserved");
assert_eq!(id2, loaded_id2, "ID2 should be preserved");
}
#[tokio::test]
async fn test_concurrent_operations() {
use futures::future::join_all;
let temp_dir = tempdir().unwrap();
let map_path = temp_dir.path().join("id_map.json");
let service = std::sync::Arc::new(IdMappingService::new(map_path).await.unwrap());
// Create multiple tasks that attempt simultaneous access
let mut tasks = Vec::new();
for i in 0..100 {
let path = StoragePath::from_string(&format!("/test/concurrent/file{}.txt", i));
let service_clone = service.clone();
tasks.push(tokio::spawn(async move {
service_clone.get_or_create_id(&path).await
}));
}
// Wait for all to finish
let results = join_all(tasks).await;
// Verify that all succeeded
for result in results {
assert!(result.unwrap().is_ok(), "Concurrent operations should succeed");
assert!(
result.unwrap().is_ok(),
"Concurrent operations should succeed"
);
}
// Save changes
service.save_pending_changes().await.unwrap();
}
}
}
@@ -1,459 +1,484 @@
//! Image Transcoding Service - WebP On-Demand Conversion
//!
//! Automatically transcodes images to WebP format when the browser supports it,
//! reducing bandwidth by 30-50% compared to JPEG/PNG.
//!
//! Features:
//! - Detects browser WebP support via Accept header
//! - Caches transcoded versions to avoid re-conversion
//! - Supports JPEG, PNG, GIF → WebP conversion
//! - Configurable quality settings
//! - Falls back to original if conversion fails
use std::path::{Path, PathBuf};
use std::sync::Arc;
use tokio::sync::RwLock;
use tokio::fs;
use bytes::Bytes;
use lru::LruCache;
use std::num::NonZeroUsize;
use image::{ImageFormat, DynamicImage};
use async_trait::async_trait;
use crate::application::ports::transcode_ports::{
ImageTranscodePort,
OutputFormat as PortOutputFormat,
TranscodeStatsDto,
};
use crate::domain::errors::{DomainError, ErrorKind};
/// Maximum file size for transcoding (5MB - larger files stream directly)
pub const MAX_TRANSCODE_SIZE: u64 = 5 * 1024 * 1024;
/// Cache key for transcoded images
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct TranscodeKey {
file_id: String,
format: OutputFormat,
}
/// Supported output formats
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum OutputFormat {
WebP,
// Future: AVIF, JPEG-XL
}
impl OutputFormat {
pub fn extension(&self) -> &'static str {
match self {
OutputFormat::WebP => "webp",
}
}
pub fn mime_type(&self) -> &'static str {
match self {
OutputFormat::WebP => "image/webp",
}
}
}
/// Result of checking browser support
#[derive(Debug)]
pub struct BrowserCapabilities {
pub supports_webp: bool,
pub supports_avif: bool,
}
impl BrowserCapabilities {
/// Parse Accept header to determine browser image format support
pub fn from_accept_header(accept: Option<&str>) -> Self {
let accept = accept.unwrap_or("");
Self {
supports_webp: accept.contains("image/webp"),
supports_avif: accept.contains("image/avif"),
}
}
/// Get the best output format for this browser
pub fn best_format(&self) -> Option<OutputFormat> {
// WebP has best support currently
if self.supports_webp {
Some(OutputFormat::WebP)
} else {
None
}
}
}
/// Image Transcoding Service
pub struct ImageTranscodeService {
/// Cache directory for transcoded images
cache_dir: PathBuf,
/// In-memory LRU cache for hot transcoded images
memory_cache: Arc<RwLock<LruCache<TranscodeKey, Bytes>>>,
/// Maximum memory cache size in bytes
max_memory_bytes: usize,
/// Current memory usage
current_memory_bytes: Arc<RwLock<usize>>,
/// Statistics
stats: Arc<RwLock<TranscodeStats>>,
}
/// Transcoding statistics
#[derive(Debug, Default, Clone)]
pub struct TranscodeStats {
pub cache_hits: u64,
pub disk_hits: u64,
pub transcodes: u64,
pub bytes_saved: u64,
pub transcode_errors: u64,
}
impl ImageTranscodeService {
/// Create new transcoding service
pub fn new(storage_root: &Path, max_cache_entries: usize, max_memory_bytes: usize) -> Self {
let cache_dir = storage_root.join(".transcoded");
Self {
cache_dir,
memory_cache: Arc::new(RwLock::new(LruCache::new(
NonZeroUsize::new(max_cache_entries).unwrap_or(NonZeroUsize::new(1000).unwrap())
))),
max_memory_bytes,
current_memory_bytes: Arc::new(RwLock::new(0)),
stats: Arc::new(RwLock::new(TranscodeStats::default())),
}
}
/// Initialize the service (create cache directories)
pub async fn initialize(&self) -> std::io::Result<()> {
fs::create_dir_all(&self.cache_dir).await?;
fs::create_dir_all(self.cache_dir.join("webp")).await?;
tracing::info!("🖼️ Image transcode service initialized at {:?}", self.cache_dir);
Ok(())
}
/// Check if a mime type can be transcoded
pub fn can_transcode(mime_type: &str) -> bool {
matches!(
mime_type,
"image/jpeg" | "image/jpg" | "image/png" | "image/gif"
)
}
/// Check if transcoding should be attempted based on file size and type
pub fn should_transcode(mime_type: &str, file_size: u64) -> bool {
Self::can_transcode(mime_type) && file_size <= MAX_TRANSCODE_SIZE
}
/// Get transcoded version of an image
/// Returns (content, mime_type, was_transcoded)
pub async fn get_transcoded(
&self,
file_id: &str,
original_content: &[u8],
original_mime: &str,
target_format: OutputFormat,
) -> Result<(Bytes, String, bool), String> {
let key = TranscodeKey {
file_id: file_id.to_string(),
format: target_format,
};
// Check memory cache first
{
let mut cache = self.memory_cache.write().await;
if let Some(cached) = cache.get(&key) {
let mut stats = self.stats.write().await;
stats.cache_hits += 1;
tracing::debug!("🔥 Transcode memory cache HIT: {}", file_id);
return Ok((cached.clone(), target_format.mime_type().to_string(), true));
}
}
// Check disk cache
let cache_path = self.get_cache_path(file_id, target_format);
if cache_path.exists() {
match fs::read(&cache_path).await {
Ok(data) => {
let content = Bytes::from(data);
// Store in memory cache
self.cache_in_memory(&key, content.clone()).await;
let mut stats = self.stats.write().await;
stats.disk_hits += 1;
tracing::debug!("💾 Transcode disk cache HIT: {}", file_id);
return Ok((content, target_format.mime_type().to_string(), true));
},
Err(e) => {
tracing::warn!("Failed to read cached transcode: {}", e);
}
}
}
// Need to transcode
let transcoded = self.transcode_image(original_content, original_mime, target_format)?;
let transcoded_bytes = Bytes::from(transcoded.clone());
// Calculate savings
let original_size = original_content.len();
let transcoded_size = transcoded_bytes.len();
let saved = original_size.saturating_sub(transcoded_size);
// Only use transcoded if it's actually smaller
if transcoded_size >= original_size {
tracing::debug!(
"⚠️ Transcode not beneficial for {}: {} -> {} bytes",
file_id, original_size, transcoded_size
);
return Ok((Bytes::from(original_content.to_vec()), original_mime.to_string(), false));
}
// Save to disk cache (async, don't wait)
let cache_path_clone = cache_path.clone();
let transcoded_clone = transcoded.clone();
tokio::spawn(async move {
if let Some(parent) = cache_path_clone.parent() {
let _ = fs::create_dir_all(parent).await;
}
if let Err(e) = fs::write(&cache_path_clone, &transcoded_clone).await {
tracing::warn!("Failed to cache transcoded image: {}", e);
}
});
// Store in memory cache
self.cache_in_memory(&key, transcoded_bytes.clone()).await;
// Update stats
{
let mut stats = self.stats.write().await;
stats.transcodes += 1;
stats.bytes_saved += saved as u64;
}
tracing::info!(
"✨ Transcoded {}: {} -> {} bytes ({:.1}% smaller)",
file_id,
original_size,
transcoded_size,
(1.0 - transcoded_size as f64 / original_size as f64) * 100.0
);
Ok((transcoded_bytes, target_format.mime_type().to_string(), true))
}
/// Perform actual image transcoding
fn transcode_image(
&self,
content: &[u8],
original_mime: &str,
target_format: OutputFormat,
) -> Result<Vec<u8>, String> {
// Determine input format
let input_format = match original_mime {
"image/jpeg" | "image/jpg" => ImageFormat::Jpeg,
"image/png" => ImageFormat::Png,
"image/gif" => ImageFormat::Gif,
_ => return Err(format!("Unsupported input format: {}", original_mime)),
};
// Load image
let img = image::load_from_memory_with_format(content, input_format)
.map_err(|e| format!("Failed to decode image: {}", e))?;
// Encode to target format
match target_format {
OutputFormat::WebP => self.encode_webp(&img),
}
}
/// Encode image to WebP
fn encode_webp(&self, img: &DynamicImage) -> Result<Vec<u8>, String> {
let mut buffer = Vec::new();
let mut cursor = std::io::Cursor::new(&mut buffer);
// Use image crate's WebP encoder
img.write_to(&mut cursor, ImageFormat::WebP)
.map_err(|e| format!("Failed to encode WebP: {}", e))?;
Ok(buffer)
}
/// Get path for cached transcoded file
fn get_cache_path(&self, file_id: &str, format: OutputFormat) -> PathBuf {
self.cache_dir
.join(format.extension())
.join(format!("{}.{}", file_id, format.extension()))
}
/// Store transcoded image in memory cache
async fn cache_in_memory(&self, key: &TranscodeKey, content: Bytes) {
let size = content.len();
let mut current = self.current_memory_bytes.write().await;
// Evict if needed
while *current + size > self.max_memory_bytes {
let mut cache = self.memory_cache.write().await;
if let Some((_, evicted)) = cache.pop_lru() {
*current = current.saturating_sub(evicted.len());
} else {
break;
}
}
// Add to cache
if *current + size <= self.max_memory_bytes {
let mut cache = self.memory_cache.write().await;
cache.put(key.clone(), content);
*current += size;
}
}
/// Invalidate cached transcodes for a file
pub async fn invalidate(&self, file_id: &str) {
// Remove from memory cache
{
let mut cache = self.memory_cache.write().await;
let key = TranscodeKey {
file_id: file_id.to_string(),
format: OutputFormat::WebP,
};
if let Some(removed) = cache.pop(&key) {
let mut current = self.current_memory_bytes.write().await;
*current = current.saturating_sub(removed.len());
}
}
// Remove disk cache
let cache_path = self.get_cache_path(file_id, OutputFormat::WebP);
let _ = fs::remove_file(&cache_path).await;
}
/// Get transcoding statistics
pub async fn get_stats(&self) -> TranscodeStats {
self.stats.read().await.clone()
}
/// Clear all caches
pub async fn clear_cache(&self) -> std::io::Result<()> {
// Clear memory
{
let mut cache = self.memory_cache.write().await;
cache.clear();
let mut current = self.current_memory_bytes.write().await;
*current = 0;
}
// Clear disk
if self.cache_dir.exists() {
fs::remove_dir_all(&self.cache_dir).await?;
fs::create_dir_all(&self.cache_dir).await?;
fs::create_dir_all(self.cache_dir.join("webp")).await?;
}
Ok(())
}
}
// ─── Port implementation ─────────────────────────────────────────────────────
/// Convert port OutputFormat to infra OutputFormat.
impl From<PortOutputFormat> for OutputFormat {
fn from(fmt: PortOutputFormat) -> Self {
match fmt {
PortOutputFormat::WebP => OutputFormat::WebP,
}
}
}
#[async_trait]
impl ImageTranscodePort for ImageTranscodeService {
fn can_transcode(&self, mime_type: &str) -> bool {
ImageTranscodeService::can_transcode(mime_type)
}
fn should_transcode(&self, mime_type: &str, file_size: u64) -> bool {
ImageTranscodeService::should_transcode(mime_type, file_size)
}
async fn get_transcoded(
&self,
file_id: &str,
original_content: &[u8],
original_mime: &str,
target_format: PortOutputFormat,
) -> Result<(Bytes, String, bool), DomainError> {
self.get_transcoded(file_id, original_content, original_mime, target_format.into())
.await
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ImageTranscode", e))
}
async fn invalidate(&self, file_id: &str) {
self.invalidate(file_id).await
}
async fn get_stats(&self) -> TranscodeStatsDto {
let stats = self.get_stats().await;
TranscodeStatsDto {
cache_hits: stats.cache_hits,
disk_hits: stats.disk_hits,
transcodes: stats.transcodes,
bytes_saved: stats.bytes_saved,
transcode_errors: stats.transcode_errors,
}
}
async fn clear_cache(&self) -> Result<(), DomainError> {
self.clear_cache().await.map_err(DomainError::from)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_browser_capabilities() {
// Chrome/Firefox with WebP support
let caps = BrowserCapabilities::from_accept_header(
Some("image/avif,image/webp,image/apng,image/svg+xml,image/*,*/*;q=0.8")
);
assert!(caps.supports_webp);
assert!(caps.supports_avif);
// Safari without WebP (old)
let caps = BrowserCapabilities::from_accept_header(
Some("image/png,image/svg+xml,image/*;q=0.8,*/*;q=0.5")
);
assert!(!caps.supports_webp);
// No header
let caps = BrowserCapabilities::from_accept_header(None);
assert!(!caps.supports_webp);
}
#[test]
fn test_can_transcode() {
assert!(ImageTranscodeService::can_transcode("image/jpeg"));
assert!(ImageTranscodeService::can_transcode("image/png"));
assert!(ImageTranscodeService::can_transcode("image/gif"));
assert!(!ImageTranscodeService::can_transcode("image/webp"));
assert!(!ImageTranscodeService::can_transcode("image/svg+xml"));
assert!(!ImageTranscodeService::can_transcode("application/pdf"));
}
#[test]
fn test_should_transcode() {
// Small JPEG - yes
assert!(ImageTranscodeService::should_transcode("image/jpeg", 1024 * 1024));
// Large JPEG - no (too big)
assert!(!ImageTranscodeService::should_transcode("image/jpeg", 10 * 1024 * 1024));
// WebP - no (already optimal)
assert!(!ImageTranscodeService::should_transcode("image/webp", 1024 * 1024));
}
}
//! Image Transcoding Service - WebP On-Demand Conversion
//!
//! Automatically transcodes images to WebP format when the browser supports it,
//! reducing bandwidth by 30-50% compared to JPEG/PNG.
//!
//! Features:
//! - Detects browser WebP support via Accept header
//! - Caches transcoded versions to avoid re-conversion
//! - Supports JPEG, PNG, GIF → WebP conversion
//! - Configurable quality settings
//! - Falls back to original if conversion fails
use async_trait::async_trait;
use bytes::Bytes;
use image::{DynamicImage, ImageFormat};
use lru::LruCache;
use std::num::NonZeroUsize;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use tokio::fs;
use tokio::sync::RwLock;
use crate::application::ports::transcode_ports::{
ImageTranscodePort, OutputFormat as PortOutputFormat, TranscodeStatsDto,
};
use crate::domain::errors::{DomainError, ErrorKind};
/// Maximum file size for transcoding (5MB - larger files stream directly)
pub const MAX_TRANSCODE_SIZE: u64 = 5 * 1024 * 1024;
/// Cache key for transcoded images
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct TranscodeKey {
file_id: String,
format: OutputFormat,
}
/// Supported output formats
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum OutputFormat {
WebP,
// Future: AVIF, JPEG-XL
}
impl OutputFormat {
pub fn extension(&self) -> &'static str {
match self {
OutputFormat::WebP => "webp",
}
}
pub fn mime_type(&self) -> &'static str {
match self {
OutputFormat::WebP => "image/webp",
}
}
}
/// Result of checking browser support
#[derive(Debug)]
pub struct BrowserCapabilities {
pub supports_webp: bool,
pub supports_avif: bool,
}
impl BrowserCapabilities {
/// Parse Accept header to determine browser image format support
pub fn from_accept_header(accept: Option<&str>) -> Self {
let accept = accept.unwrap_or("");
Self {
supports_webp: accept.contains("image/webp"),
supports_avif: accept.contains("image/avif"),
}
}
/// Get the best output format for this browser
pub fn best_format(&self) -> Option<OutputFormat> {
// WebP has best support currently
if self.supports_webp {
Some(OutputFormat::WebP)
} else {
None
}
}
}
/// Image Transcoding Service
pub struct ImageTranscodeService {
/// Cache directory for transcoded images
cache_dir: PathBuf,
/// In-memory LRU cache for hot transcoded images
memory_cache: Arc<RwLock<LruCache<TranscodeKey, Bytes>>>,
/// Maximum memory cache size in bytes
max_memory_bytes: usize,
/// Current memory usage
current_memory_bytes: Arc<RwLock<usize>>,
/// Statistics
stats: Arc<RwLock<TranscodeStats>>,
}
/// Transcoding statistics
#[derive(Debug, Default, Clone)]
pub struct TranscodeStats {
pub cache_hits: u64,
pub disk_hits: u64,
pub transcodes: u64,
pub bytes_saved: u64,
pub transcode_errors: u64,
}
impl ImageTranscodeService {
/// Create new transcoding service
pub fn new(storage_root: &Path, max_cache_entries: usize, max_memory_bytes: usize) -> Self {
let cache_dir = storage_root.join(".transcoded");
Self {
cache_dir,
memory_cache: Arc::new(RwLock::new(LruCache::new(
NonZeroUsize::new(max_cache_entries).unwrap_or(NonZeroUsize::new(1000).unwrap()),
))),
max_memory_bytes,
current_memory_bytes: Arc::new(RwLock::new(0)),
stats: Arc::new(RwLock::new(TranscodeStats::default())),
}
}
/// Initialize the service (create cache directories)
pub async fn initialize(&self) -> std::io::Result<()> {
fs::create_dir_all(&self.cache_dir).await?;
fs::create_dir_all(self.cache_dir.join("webp")).await?;
tracing::info!(
"🖼️ Image transcode service initialized at {:?}",
self.cache_dir
);
Ok(())
}
/// Check if a mime type can be transcoded
pub fn can_transcode(mime_type: &str) -> bool {
matches!(
mime_type,
"image/jpeg" | "image/jpg" | "image/png" | "image/gif"
)
}
/// Check if transcoding should be attempted based on file size and type
pub fn should_transcode(mime_type: &str, file_size: u64) -> bool {
Self::can_transcode(mime_type) && file_size <= MAX_TRANSCODE_SIZE
}
/// Get transcoded version of an image
/// Returns (content, mime_type, was_transcoded)
pub async fn get_transcoded(
&self,
file_id: &str,
original_content: &[u8],
original_mime: &str,
target_format: OutputFormat,
) -> Result<(Bytes, String, bool), String> {
let key = TranscodeKey {
file_id: file_id.to_string(),
format: target_format,
};
// Check memory cache first
{
let mut cache = self.memory_cache.write().await;
if let Some(cached) = cache.get(&key) {
let mut stats = self.stats.write().await;
stats.cache_hits += 1;
tracing::debug!("🔥 Transcode memory cache HIT: {}", file_id);
return Ok((cached.clone(), target_format.mime_type().to_string(), true));
}
}
// Check disk cache
let cache_path = self.get_cache_path(file_id, target_format);
if cache_path.exists() {
match fs::read(&cache_path).await {
Ok(data) => {
let content = Bytes::from(data);
// Store in memory cache
self.cache_in_memory(&key, content.clone()).await;
let mut stats = self.stats.write().await;
stats.disk_hits += 1;
tracing::debug!("💾 Transcode disk cache HIT: {}", file_id);
return Ok((content, target_format.mime_type().to_string(), true));
}
Err(e) => {
tracing::warn!("Failed to read cached transcode: {}", e);
}
}
}
// Need to transcode
let transcoded = self.transcode_image(original_content, original_mime, target_format)?;
let transcoded_bytes = Bytes::from(transcoded.clone());
// Calculate savings
let original_size = original_content.len();
let transcoded_size = transcoded_bytes.len();
let saved = original_size.saturating_sub(transcoded_size);
// Only use transcoded if it's actually smaller
if transcoded_size >= original_size {
tracing::debug!(
"⚠️ Transcode not beneficial for {}: {} -> {} bytes",
file_id,
original_size,
transcoded_size
);
return Ok((
Bytes::from(original_content.to_vec()),
original_mime.to_string(),
false,
));
}
// Save to disk cache (async, don't wait)
let cache_path_clone = cache_path.clone();
let transcoded_clone = transcoded.clone();
tokio::spawn(async move {
if let Some(parent) = cache_path_clone.parent() {
let _ = fs::create_dir_all(parent).await;
}
if let Err(e) = fs::write(&cache_path_clone, &transcoded_clone).await {
tracing::warn!("Failed to cache transcoded image: {}", e);
}
});
// Store in memory cache
self.cache_in_memory(&key, transcoded_bytes.clone()).await;
// Update stats
{
let mut stats = self.stats.write().await;
stats.transcodes += 1;
stats.bytes_saved += saved as u64;
}
tracing::info!(
"✨ Transcoded {}: {} -> {} bytes ({:.1}% smaller)",
file_id,
original_size,
transcoded_size,
(1.0 - transcoded_size as f64 / original_size as f64) * 100.0
);
Ok((
transcoded_bytes,
target_format.mime_type().to_string(),
true,
))
}
/// Perform actual image transcoding
fn transcode_image(
&self,
content: &[u8],
original_mime: &str,
target_format: OutputFormat,
) -> Result<Vec<u8>, String> {
// Determine input format
let input_format = match original_mime {
"image/jpeg" | "image/jpg" => ImageFormat::Jpeg,
"image/png" => ImageFormat::Png,
"image/gif" => ImageFormat::Gif,
_ => return Err(format!("Unsupported input format: {}", original_mime)),
};
// Load image
let img = image::load_from_memory_with_format(content, input_format)
.map_err(|e| format!("Failed to decode image: {}", e))?;
// Encode to target format
match target_format {
OutputFormat::WebP => self.encode_webp(&img),
}
}
/// Encode image to WebP
fn encode_webp(&self, img: &DynamicImage) -> Result<Vec<u8>, String> {
let mut buffer = Vec::new();
let mut cursor = std::io::Cursor::new(&mut buffer);
// Use image crate's WebP encoder
img.write_to(&mut cursor, ImageFormat::WebP)
.map_err(|e| format!("Failed to encode WebP: {}", e))?;
Ok(buffer)
}
/// Get path for cached transcoded file
fn get_cache_path(&self, file_id: &str, format: OutputFormat) -> PathBuf {
self.cache_dir
.join(format.extension())
.join(format!("{}.{}", file_id, format.extension()))
}
/// Store transcoded image in memory cache
async fn cache_in_memory(&self, key: &TranscodeKey, content: Bytes) {
let size = content.len();
let mut current = self.current_memory_bytes.write().await;
// Evict if needed
while *current + size > self.max_memory_bytes {
let mut cache = self.memory_cache.write().await;
if let Some((_, evicted)) = cache.pop_lru() {
*current = current.saturating_sub(evicted.len());
} else {
break;
}
}
// Add to cache
if *current + size <= self.max_memory_bytes {
let mut cache = self.memory_cache.write().await;
cache.put(key.clone(), content);
*current += size;
}
}
/// Invalidate cached transcodes for a file
pub async fn invalidate(&self, file_id: &str) {
// Remove from memory cache
{
let mut cache = self.memory_cache.write().await;
let key = TranscodeKey {
file_id: file_id.to_string(),
format: OutputFormat::WebP,
};
if let Some(removed) = cache.pop(&key) {
let mut current = self.current_memory_bytes.write().await;
*current = current.saturating_sub(removed.len());
}
}
// Remove disk cache
let cache_path = self.get_cache_path(file_id, OutputFormat::WebP);
let _ = fs::remove_file(&cache_path).await;
}
/// Get transcoding statistics
pub async fn get_stats(&self) -> TranscodeStats {
self.stats.read().await.clone()
}
/// Clear all caches
pub async fn clear_cache(&self) -> std::io::Result<()> {
// Clear memory
{
let mut cache = self.memory_cache.write().await;
cache.clear();
let mut current = self.current_memory_bytes.write().await;
*current = 0;
}
// Clear disk
if self.cache_dir.exists() {
fs::remove_dir_all(&self.cache_dir).await?;
fs::create_dir_all(&self.cache_dir).await?;
fs::create_dir_all(self.cache_dir.join("webp")).await?;
}
Ok(())
}
}
// ─── Port implementation ─────────────────────────────────────────────────────
/// Convert port OutputFormat to infra OutputFormat.
impl From<PortOutputFormat> for OutputFormat {
fn from(fmt: PortOutputFormat) -> Self {
match fmt {
PortOutputFormat::WebP => OutputFormat::WebP,
}
}
}
#[async_trait]
impl ImageTranscodePort for ImageTranscodeService {
fn can_transcode(&self, mime_type: &str) -> bool {
ImageTranscodeService::can_transcode(mime_type)
}
fn should_transcode(&self, mime_type: &str, file_size: u64) -> bool {
ImageTranscodeService::should_transcode(mime_type, file_size)
}
async fn get_transcoded(
&self,
file_id: &str,
original_content: &[u8],
original_mime: &str,
target_format: PortOutputFormat,
) -> Result<(Bytes, String, bool), DomainError> {
self.get_transcoded(
file_id,
original_content,
original_mime,
target_format.into(),
)
.await
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ImageTranscode", e))
}
async fn invalidate(&self, file_id: &str) {
self.invalidate(file_id).await
}
async fn get_stats(&self) -> TranscodeStatsDto {
let stats = self.get_stats().await;
TranscodeStatsDto {
cache_hits: stats.cache_hits,
disk_hits: stats.disk_hits,
transcodes: stats.transcodes,
bytes_saved: stats.bytes_saved,
transcode_errors: stats.transcode_errors,
}
}
async fn clear_cache(&self) -> Result<(), DomainError> {
self.clear_cache().await.map_err(DomainError::from)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_browser_capabilities() {
// Chrome/Firefox with WebP support
let caps = BrowserCapabilities::from_accept_header(Some(
"image/avif,image/webp,image/apng,image/svg+xml,image/*,*/*;q=0.8",
));
assert!(caps.supports_webp);
assert!(caps.supports_avif);
// Safari without WebP (old)
let caps = BrowserCapabilities::from_accept_header(Some(
"image/png,image/svg+xml,image/*;q=0.8,*/*;q=0.5",
));
assert!(!caps.supports_webp);
// No header
let caps = BrowserCapabilities::from_accept_header(None);
assert!(!caps.supports_webp);
}
#[test]
fn test_can_transcode() {
assert!(ImageTranscodeService::can_transcode("image/jpeg"));
assert!(ImageTranscodeService::can_transcode("image/png"));
assert!(ImageTranscodeService::can_transcode("image/gif"));
assert!(!ImageTranscodeService::can_transcode("image/webp"));
assert!(!ImageTranscodeService::can_transcode("image/svg+xml"));
assert!(!ImageTranscodeService::can_transcode("application/pdf"));
}
#[test]
fn test_should_transcode() {
// Small JPEG - yes
assert!(ImageTranscodeService::should_transcode(
"image/jpeg",
1024 * 1024
));
// Large JPEG - no (too big)
assert!(!ImageTranscodeService::should_transcode(
"image/jpeg",
10 * 1024 * 1024
));
// WebP - no (already optimal)
assert!(!ImageTranscodeService::should_transcode(
"image/webp",
1024 * 1024
));
}
}
+221 -210
View File
@@ -1,210 +1,221 @@
//! JWT-based token service implementation.
//!
//! This module provides JWT token generation and validation functionality,
//! implementing the TokenServicePort trait defined in the application layer.
use jsonwebtoken::{encode, decode, Header, Validation, EncodingKey, DecodingKey, Algorithm};
use serde::{Serialize, Deserialize};
use uuid::Uuid;
use chrono::Utc;
use crate::application::ports::auth_ports::{TokenServicePort, TokenClaims};
use crate::domain::entities::user::User;
use crate::common::errors::{DomainError, ErrorKind};
/// Internal JWT claims structure for serialization.
/// This is the actual JWT payload structure used by jsonwebtoken crate.
#[derive(Debug, Serialize, Deserialize)]
struct JwtClaims {
/// Subject identifier - contains the user ID
pub sub: String,
/// Expiration timestamp (seconds since Unix epoch)
pub exp: i64,
/// Issued at timestamp (seconds since Unix epoch)
pub iat: i64,
/// JWT unique ID for token tracking and revocation
pub jti: String,
/// Username for display and identification purposes
pub username: String,
/// User email for communication and identification
pub email: String,
/// User role for authorization checks
pub role: String,
}
impl From<JwtClaims> for TokenClaims {
fn from(claims: JwtClaims) -> Self {
TokenClaims {
sub: claims.sub,
exp: claims.exp,
iat: claims.iat,
jti: claims.jti,
username: claims.username,
email: claims.email,
role: claims.role,
}
}
}
/// JWT-based implementation of the TokenServicePort.
///
/// This service handles JWT token generation and validation for user authentication.
/// It uses HS256 algorithm for signing tokens.
pub struct JwtTokenService {
/// Secret key used for signing JWT tokens
jwt_secret: String,
/// Expiration time for access tokens in seconds
access_token_expiry: i64,
/// Expiration time for refresh tokens in seconds
refresh_token_expiry: i64,
}
impl JwtTokenService {
/// Create a new JwtTokenService with the specified configuration.
///
/// # Arguments
/// * `jwt_secret` - Secret key for signing tokens (should be at least 32 bytes)
/// * `access_token_expiry_secs` - Lifetime of access tokens in seconds
/// * `refresh_token_expiry_secs` - Lifetime of refresh tokens in seconds
pub fn new(jwt_secret: String, access_token_expiry_secs: i64, refresh_token_expiry_secs: i64) -> Self {
Self {
jwt_secret,
access_token_expiry: access_token_expiry_secs,
refresh_token_expiry: refresh_token_expiry_secs,
}
}
}
impl TokenServicePort for JwtTokenService {
fn generate_access_token(&self, user: &User) -> Result<String, DomainError> {
let now = Utc::now().timestamp();
// Log information for debugging
tracing::debug!(
"Generating token for user: {}, id: {}, role: {}",
user.username(),
user.id(),
user.role()
);
let claims = JwtClaims {
sub: user.id().to_string(),
exp: now + self.access_token_expiry,
iat: now,
jti: Uuid::new_v4().to_string(),
username: user.username().to_string(),
email: user.email().to_string(),
role: format!("{}", user.role()),
};
// Log JWT claims for debugging
tracing::debug!("JWT claims: sub={}, exp={}, iat={}", claims.sub, claims.exp, claims.iat);
encode(
&Header::default(),
&claims,
&EncodingKey::from_secret(self.jwt_secret.as_bytes())
)
.map_err(|e| {
tracing::error!("Error generating token: {}", e);
DomainError::new(
ErrorKind::InternalError,
"TokenService",
format!("Error generating token: {}", e)
)
})
}
fn validate_token(&self, token: &str) -> Result<TokenClaims, DomainError> {
let validation = Validation::new(Algorithm::HS256);
let token_data = decode::<JwtClaims>(
token,
&DecodingKey::from_secret(self.jwt_secret.as_bytes()),
&validation
)
.map_err(|e| {
match e.kind() {
jsonwebtoken::errors::ErrorKind::ExpiredSignature => {
DomainError::new(ErrorKind::AccessDenied, "TokenService", "Token expired")
},
_ => DomainError::new(
ErrorKind::AccessDenied,
"TokenService",
format!("Invalid token: {}", e)
),
}
})?;
Ok(token_data.claims.into())
}
fn generate_refresh_token(&self) -> String {
Uuid::new_v4().to_string()
}
fn refresh_token_expiry_secs(&self) -> i64 {
self.refresh_token_expiry
}
fn refresh_token_expiry_days(&self) -> i64 {
self.refresh_token_expiry / (24 * 3600)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::domain::entities::user::{User, UserRole};
fn create_test_user() -> User {
User::from_data(
"test-user-id".to_string(),
"testuser".to_string(),
"test@example.com".to_string(),
"hashed_password".to_string(),
UserRole::User,
1024 * 1024 * 1024, // 1GB
0,
chrono::Utc::now(),
chrono::Utc::now(),
None,
true,
)
}
#[test]
fn test_generate_and_validate_token() {
let service = JwtTokenService::new(
"test_secret_key_at_least_32_bytes_long".to_string(),
3600, // 1 hour
86400, // 1 day
);
let user = create_test_user();
let token = service.generate_access_token(&user).expect("Should generate token");
let claims = service.validate_token(&token).expect("Should validate token");
assert_eq!(claims.sub, user.id());
assert_eq!(claims.username, user.username());
assert_eq!(claims.email, user.email());
}
#[test]
fn test_refresh_token_is_unique() {
let service = JwtTokenService::new("secret".to_string(), 3600, 86400);
let token1 = service.generate_refresh_token();
let token2 = service.generate_refresh_token();
assert_ne!(token1, token2);
}
#[test]
fn test_invalid_token() {
let service = JwtTokenService::new("secret".to_string(), 3600, 86400);
let result = service.validate_token("invalid_token");
assert!(result.is_err());
}
}
//! JWT-based token service implementation.
//!
//! This module provides JWT token generation and validation functionality,
//! implementing the TokenServicePort trait defined in the application layer.
use chrono::Utc;
use jsonwebtoken::{Algorithm, DecodingKey, EncodingKey, Header, Validation, decode, encode};
use serde::{Deserialize, Serialize};
use uuid::Uuid;
use crate::application::ports::auth_ports::{TokenClaims, TokenServicePort};
use crate::common::errors::{DomainError, ErrorKind};
use crate::domain::entities::user::User;
/// Internal JWT claims structure for serialization.
/// This is the actual JWT payload structure used by jsonwebtoken crate.
#[derive(Debug, Serialize, Deserialize)]
struct JwtClaims {
/// Subject identifier - contains the user ID
pub sub: String,
/// Expiration timestamp (seconds since Unix epoch)
pub exp: i64,
/// Issued at timestamp (seconds since Unix epoch)
pub iat: i64,
/// JWT unique ID for token tracking and revocation
pub jti: String,
/// Username for display and identification purposes
pub username: String,
/// User email for communication and identification
pub email: String,
/// User role for authorization checks
pub role: String,
}
impl From<JwtClaims> for TokenClaims {
fn from(claims: JwtClaims) -> Self {
TokenClaims {
sub: claims.sub,
exp: claims.exp,
iat: claims.iat,
jti: claims.jti,
username: claims.username,
email: claims.email,
role: claims.role,
}
}
}
/// JWT-based implementation of the TokenServicePort.
///
/// This service handles JWT token generation and validation for user authentication.
/// It uses HS256 algorithm for signing tokens.
pub struct JwtTokenService {
/// Secret key used for signing JWT tokens
jwt_secret: String,
/// Expiration time for access tokens in seconds
access_token_expiry: i64,
/// Expiration time for refresh tokens in seconds
refresh_token_expiry: i64,
}
impl JwtTokenService {
/// Create a new JwtTokenService with the specified configuration.
///
/// # Arguments
/// * `jwt_secret` - Secret key for signing tokens (should be at least 32 bytes)
/// * `access_token_expiry_secs` - Lifetime of access tokens in seconds
/// * `refresh_token_expiry_secs` - Lifetime of refresh tokens in seconds
pub fn new(
jwt_secret: String,
access_token_expiry_secs: i64,
refresh_token_expiry_secs: i64,
) -> Self {
Self {
jwt_secret,
access_token_expiry: access_token_expiry_secs,
refresh_token_expiry: refresh_token_expiry_secs,
}
}
}
impl TokenServicePort for JwtTokenService {
fn generate_access_token(&self, user: &User) -> Result<String, DomainError> {
let now = Utc::now().timestamp();
// Log information for debugging
tracing::debug!(
"Generating token for user: {}, id: {}, role: {}",
user.username(),
user.id(),
user.role()
);
let claims = JwtClaims {
sub: user.id().to_string(),
exp: now + self.access_token_expiry,
iat: now,
jti: Uuid::new_v4().to_string(),
username: user.username().to_string(),
email: user.email().to_string(),
role: format!("{}", user.role()),
};
// Log JWT claims for debugging
tracing::debug!(
"JWT claims: sub={}, exp={}, iat={}",
claims.sub,
claims.exp,
claims.iat
);
encode(
&Header::default(),
&claims,
&EncodingKey::from_secret(self.jwt_secret.as_bytes()),
)
.map_err(|e| {
tracing::error!("Error generating token: {}", e);
DomainError::new(
ErrorKind::InternalError,
"TokenService",
format!("Error generating token: {}", e),
)
})
}
fn validate_token(&self, token: &str) -> Result<TokenClaims, DomainError> {
let validation = Validation::new(Algorithm::HS256);
let token_data = decode::<JwtClaims>(
token,
&DecodingKey::from_secret(self.jwt_secret.as_bytes()),
&validation,
)
.map_err(|e| match e.kind() {
jsonwebtoken::errors::ErrorKind::ExpiredSignature => {
DomainError::new(ErrorKind::AccessDenied, "TokenService", "Token expired")
}
_ => DomainError::new(
ErrorKind::AccessDenied,
"TokenService",
format!("Invalid token: {}", e),
),
})?;
Ok(token_data.claims.into())
}
fn generate_refresh_token(&self) -> String {
Uuid::new_v4().to_string()
}
fn refresh_token_expiry_secs(&self) -> i64 {
self.refresh_token_expiry
}
fn refresh_token_expiry_days(&self) -> i64 {
self.refresh_token_expiry / (24 * 3600)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::domain::entities::user::{User, UserRole};
fn create_test_user() -> User {
User::from_data(
"test-user-id".to_string(),
"testuser".to_string(),
"test@example.com".to_string(),
"hashed_password".to_string(),
UserRole::User,
1024 * 1024 * 1024, // 1GB
0,
chrono::Utc::now(),
chrono::Utc::now(),
None,
true,
)
}
#[test]
fn test_generate_and_validate_token() {
let service = JwtTokenService::new(
"test_secret_key_at_least_32_bytes_long".to_string(),
3600, // 1 hour
86400, // 1 day
);
let user = create_test_user();
let token = service
.generate_access_token(&user)
.expect("Should generate token");
let claims = service
.validate_token(&token)
.expect("Should validate token");
assert_eq!(claims.sub, user.id());
assert_eq!(claims.username, user.username());
assert_eq!(claims.email, user.email());
}
#[test]
fn test_refresh_token_is_unique() {
let service = JwtTokenService::new("secret".to_string(), 3600, 86400);
let token1 = service.generate_refresh_token();
let token2 = service.generate_refresh_token();
assert_ne!(token1, token2);
}
#[test]
fn test_invalid_token() {
let service = JwtTokenService::new("secret".to_string(), 3600, 86400);
let result = service.validate_token("invalid_token");
assert!(result.is_err());
}
}
+15 -15
View File
@@ -1,19 +1,19 @@
pub mod buffer_pool;
pub mod chunked_upload_service;
pub mod compression_service;
pub mod dedup_service;
pub mod file_content_cache;
pub mod file_metadata_cache;
pub mod file_system_i18n_service;
pub mod file_system_utils;
pub mod id_mapping_service;
pub mod id_mapping_optimizer;
pub mod file_metadata_cache;
pub mod file_content_cache;
pub mod compression_service;
pub mod buffer_pool;
pub mod trash_cleanup_service;
pub mod zip_service;
pub mod path_service;
pub mod password_hasher;
pub mod jwt_service;
pub mod thumbnail_service;
pub mod write_behind_cache;
pub mod chunked_upload_service;
pub mod id_mapping_service;
pub mod image_transcode_service;
pub mod dedup_service;
pub mod oidc_service;
pub mod jwt_service;
pub mod oidc_service;
pub mod password_hasher;
pub mod path_service;
pub mod thumbnail_service;
pub mod trash_cleanup_service;
pub mod write_behind_cache;
pub mod zip_service;
+531 -449
View File
@@ -1,449 +1,531 @@
//! OpenID Connect (OIDC) service implementation.
//!
//! Handles OIDC discovery, authorization URL generation, code exchange,
//! ID token validation (RS256 via JWKS), and UserInfo fetching.
//! Compatible with Authentik, Keycloak, and any standard OIDC provider.
use std::sync::RwLock;
use async_trait::async_trait;
use serde::Deserialize;
use crate::application::ports::auth_ports::{OidcServicePort, OidcTokenSet, OidcIdClaims};
use crate::common::config::OidcConfig;
use crate::common::errors::{DomainError, ErrorKind};
// ============================================================================
// OIDC Discovery Document
// ============================================================================
#[derive(Debug, Clone, Deserialize)]
struct OidcDiscovery {
issuer: String,
authorization_endpoint: String,
token_endpoint: String,
userinfo_endpoint: Option<String>,
jwks_uri: String,
}
// ============================================================================
// JWKS structures for RS256 validation
// ============================================================================
#[derive(Debug, Clone, Deserialize)]
struct JwksDocument {
keys: Vec<JwkKey>,
}
#[derive(Debug, Clone, Deserialize)]
struct JwkKey {
kty: String,
#[serde(rename = "use")]
key_use: Option<String>,
kid: Option<String>,
alg: Option<String>,
n: Option<String>, // RSA modulus (base64url)
e: Option<String>, // RSA exponent (base64url)
}
// ============================================================================
// Token exchange response
// ============================================================================
#[derive(Debug, Deserialize)]
struct TokenResponse {
access_token: String,
id_token: Option<String>,
refresh_token: Option<String>,
#[allow(dead_code)]
token_type: Option<String>,
#[allow(dead_code)]
expires_in: Option<i64>,
}
// ============================================================================
// ID token claims (standard OIDC)
// ============================================================================
#[derive(Debug, Deserialize)]
struct IdTokenClaims {
sub: String,
email: Option<String>,
preferred_username: Option<String>,
name: Option<String>,
groups: Option<Vec<String>>,
nonce: Option<String>,
// Standard JWT fields
#[allow(dead_code)]
iss: Option<String>,
#[allow(dead_code)]
aud: Option<serde_json::Value>,
#[allow(dead_code)]
exp: Option<i64>,
#[allow(dead_code)]
iat: Option<i64>,
}
// ============================================================================
// UserInfo response
// ============================================================================
#[derive(Debug, Deserialize)]
struct UserInfoResponse {
sub: String,
email: Option<String>,
preferred_username: Option<String>,
name: Option<String>,
groups: Option<Vec<String>>,
}
// ============================================================================
// OIDC Service
// ============================================================================
pub struct OidcService {
config: OidcConfig,
http_client: reqwest::Client,
/// Cached discovery document
discovery: RwLock<Option<OidcDiscovery>>,
/// Cached JWKS
jwks: RwLock<Option<JwksDocument>>,
}
impl OidcService {
pub fn new(config: OidcConfig) -> Self {
let http_client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.expect("Failed to build HTTP client for OIDC");
Self {
config,
http_client,
discovery: RwLock::new(None),
jwks: RwLock::new(None),
}
}
/// Fetch and cache the OIDC discovery document
async fn get_discovery(&self) -> Result<OidcDiscovery, DomainError> {
// Check cache first
{
let cache = self.discovery.read().map_err(|_| DomainError::new(
ErrorKind::InternalError, "OIDC", "Lock poisoned",
))?;
if let Some(ref disc) = *cache {
return Ok(disc.clone());
}
}
// Fetch discovery document
let issuer = self.config.issuer_url.trim_end_matches('/');
let discovery_url = format!("{}/.well-known/openid-configuration", issuer);
tracing::info!("Fetching OIDC discovery from: {}", discovery_url);
let resp = self.http_client.get(&discovery_url)
.send()
.await
.map_err(|e| DomainError::new(
ErrorKind::InternalError, "OIDC",
format!("Failed to fetch OIDC discovery: {}", e),
))?;
if !resp.status().is_success() {
return Err(DomainError::new(
ErrorKind::InternalError, "OIDC",
format!("OIDC discovery returned status {}", resp.status()),
));
}
let discovery: OidcDiscovery = resp.json().await.map_err(|e| DomainError::new(
ErrorKind::InternalError, "OIDC",
format!("Failed to parse OIDC discovery: {}", e),
))?;
// Cache it
{
let mut cache = self.discovery.write().map_err(|_| DomainError::new(
ErrorKind::InternalError, "OIDC", "Lock poisoned",
))?;
*cache = Some(discovery.clone());
}
Ok(discovery)
}
/// Fetch and cache JWKS document for ID token validation
async fn get_jwks(&self) -> Result<JwksDocument, DomainError> {
// Check cache first
{
let cache = self.jwks.read().map_err(|_| DomainError::new(
ErrorKind::InternalError, "OIDC", "Lock poisoned",
))?;
if let Some(ref jwks) = *cache {
return Ok(jwks.clone());
}
}
let discovery = self.get_discovery().await?;
tracing::debug!("Fetching JWKS from: {}", discovery.jwks_uri);
let resp = self.http_client.get(&discovery.jwks_uri)
.send()
.await
.map_err(|e| DomainError::new(
ErrorKind::InternalError, "OIDC",
format!("Failed to fetch JWKS: {}", e),
))?;
let jwks: JwksDocument = resp.json().await.map_err(|e| DomainError::new(
ErrorKind::InternalError, "OIDC",
format!("Failed to parse JWKS: {}", e),
))?;
// Cache it
{
let mut cache = self.jwks.write().map_err(|_| DomainError::new(
ErrorKind::InternalError, "OIDC", "Lock poisoned",
))?;
*cache = Some(jwks.clone());
}
Ok(jwks)
}
/// Find the right RSA key from JWKS by kid header
fn find_rsa_key<'a>(jwks: &'a JwksDocument, kid: Option<&str>) -> Option<&'a JwkKey> {
jwks.keys.iter().find(|k| {
k.kty == "RSA"
&& k.key_use.as_deref() != Some("enc") // exclude encryption keys
&& (kid.is_none() || k.kid.as_deref() == kid)
})
}
/// Extract the `kid` from a JWT header without full validation
fn extract_jwt_kid(token: &str) -> Option<String> {
let parts: Vec<&str> = token.splitn(3, '.').collect();
if parts.len() < 2 {
return None;
}
use base64::Engine;
let engine = base64::engine::general_purpose::URL_SAFE_NO_PAD;
let header_bytes = engine.decode(parts[0]).ok()?;
let header: serde_json::Value = serde_json::from_slice(&header_bytes).ok()?;
header.get("kid").and_then(|v| v.as_str()).map(|s| s.to_string())
}
}
#[async_trait]
impl OidcServicePort for OidcService {
async fn get_authorize_url(&self, state: &str, nonce: &str, pkce_challenge: &str) -> Result<String, DomainError> {
// Fetch or use cached discovery to get the correct authorization_endpoint
let discovery = self.get_discovery().await?;
let auth_endpoint = discovery.authorization_endpoint;
let scopes = self.config.scopes.replace(',', " ");
let url = format!(
"{}?response_type=code&client_id={}&redirect_uri={}&scope={}&state={}&nonce={}&code_challenge={}&code_challenge_method=S256",
auth_endpoint,
urlencoding::encode(&self.config.client_id),
urlencoding::encode(&self.config.redirect_uri),
urlencoding::encode(&scopes),
urlencoding::encode(state),
urlencoding::encode(nonce),
urlencoding::encode(pkce_challenge),
);
Ok(url)
}
async fn exchange_code(&self, code: &str, pkce_verifier: &str) -> Result<OidcTokenSet, DomainError> {
let discovery = self.get_discovery().await?;
tracing::debug!("Exchanging authorization code at: {}", discovery.token_endpoint);
let resp = self.http_client.post(&discovery.token_endpoint)
.form(&[
("grant_type", "authorization_code"),
("code", code),
("redirect_uri", &self.config.redirect_uri),
("client_id", &self.config.client_id),
("client_secret", &self.config.client_secret),
("code_verifier", pkce_verifier),
])
.send()
.await
.map_err(|e| DomainError::new(
ErrorKind::InternalError, "OIDC",
format!("Token exchange failed: {}", e),
))?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
tracing::error!("OIDC token exchange error: status={}, body={}", status, body);
return Err(DomainError::new(
ErrorKind::AccessDenied, "OIDC",
format!("Token exchange failed with status {}", status),
));
}
let token_resp: TokenResponse = resp.json().await.map_err(|e| DomainError::new(
ErrorKind::InternalError, "OIDC",
format!("Failed to parse token response: {}", e),
))?;
let id_token = token_resp.id_token.ok_or_else(|| DomainError::new(
ErrorKind::InternalError, "OIDC",
"No id_token in token response",
))?;
Ok(OidcTokenSet {
access_token: token_resp.access_token,
id_token,
refresh_token: token_resp.refresh_token,
})
}
async fn validate_id_token(&self, id_token: &str, expected_nonce: Option<&str>) -> Result<OidcIdClaims, DomainError> {
let jwks = self.get_jwks().await?;
let discovery = self.get_discovery().await?;
// Extract kid from JWT header
let kid = Self::extract_jwt_kid(id_token);
// Find the matching RSA key
let jwk = Self::find_rsa_key(&jwks, kid.as_deref()).ok_or_else(|| DomainError::new(
ErrorKind::AccessDenied, "OIDC",
"No suitable RSA key found in JWKS for ID token validation",
))?;
let n = jwk.n.as_ref().ok_or_else(|| DomainError::new(
ErrorKind::InternalError, "OIDC", "JWKS key missing 'n' component",
))?;
let e = jwk.e.as_ref().ok_or_else(|| DomainError::new(
ErrorKind::InternalError, "OIDC", "JWKS key missing 'e' component",
))?;
// Build decoding key from RSA components
let decoding_key = jsonwebtoken::DecodingKey::from_rsa_components(n, e)
.map_err(|err| DomainError::new(
ErrorKind::InternalError, "OIDC",
format!("Failed to build RSA decoding key: {}", err),
))?;
// Determine algorithm from JWKS (default RS256)
let alg = match jwk.alg.as_deref() {
Some("RS384") => jsonwebtoken::Algorithm::RS384,
Some("RS512") => jsonwebtoken::Algorithm::RS512,
_ => jsonwebtoken::Algorithm::RS256,
};
// Build validation: check expiry and issuer
let mut validation = jsonwebtoken::Validation::new(alg);
validation.set_issuer(&[&discovery.issuer]);
validation.set_audience(&[&self.config.client_id]);
let token_data = jsonwebtoken::decode::<IdTokenClaims>(
id_token,
&decoding_key,
&validation,
).map_err(|e| {
tracing::warn!("OIDC ID token validation failed: {}", e);
DomainError::new(
ErrorKind::AccessDenied, "OIDC",
format!("ID token validation failed: {}", e),
)
})?;
let claims = token_data.claims;
// Verify nonce to prevent token replay attacks
if let Some(expected) = expected_nonce {
match &claims.nonce {
Some(actual) if actual == expected => { /* OK */ }
Some(actual) => {
tracing::warn!("OIDC nonce mismatch: expected={}, got={}", expected, actual);
return Err(DomainError::new(
ErrorKind::AccessDenied, "OIDC",
"ID token nonce mismatch — possible replay attack",
));
}
None => {
tracing::warn!("OIDC nonce missing from ID token (expected={})", expected);
// Some providers don't include nonce; log warning but don't fail
}
}
}
Ok(OidcIdClaims {
sub: claims.sub,
email: claims.email,
preferred_username: claims.preferred_username,
name: claims.name,
groups: claims.groups.unwrap_or_default(),
})
}
async fn fetch_user_info(&self, access_token: &str) -> Result<OidcIdClaims, DomainError> {
let discovery = self.get_discovery().await?;
let userinfo_url = discovery.userinfo_endpoint.ok_or_else(|| DomainError::new(
ErrorKind::InternalError, "OIDC",
"No userinfo_endpoint in OIDC discovery",
))?;
let resp = self.http_client.get(&userinfo_url)
.header("Authorization", format!("Bearer {}", access_token))
.send()
.await
.map_err(|e| DomainError::new(
ErrorKind::InternalError, "OIDC",
format!("UserInfo request failed: {}", e),
))?;
if !resp.status().is_success() {
return Err(DomainError::new(
ErrorKind::AccessDenied, "OIDC",
format!("UserInfo returned status {}", resp.status()),
));
}
let info: UserInfoResponse = resp.json().await.map_err(|e| DomainError::new(
ErrorKind::InternalError, "OIDC",
format!("Failed to parse UserInfo: {}", e),
))?;
Ok(OidcIdClaims {
sub: info.sub,
email: info.email,
preferred_username: info.preferred_username,
name: info.name,
groups: info.groups.unwrap_or_default(),
})
}
fn provider_name(&self) -> &str {
&self.config.provider_name
}
}
// We need urlencoding — let's use a minimal inline implementation
mod urlencoding {
pub fn encode(input: &str) -> String {
let mut result = String::with_capacity(input.len() * 3);
for byte in input.bytes() {
match byte {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
result.push(byte as char);
}
_ => {
result.push('%');
result.push_str(&format!("{:02X}", byte));
}
}
}
result
}
}
//! OpenID Connect (OIDC) service implementation.
//!
//! Handles OIDC discovery, authorization URL generation, code exchange,
//! ID token validation (RS256 via JWKS), and UserInfo fetching.
//! Compatible with Authentik, Keycloak, and any standard OIDC provider.
use async_trait::async_trait;
use serde::Deserialize;
use std::sync::RwLock;
use crate::application::ports::auth_ports::{OidcIdClaims, OidcServicePort, OidcTokenSet};
use crate::common::config::OidcConfig;
use crate::common::errors::{DomainError, ErrorKind};
// ============================================================================
// OIDC Discovery Document
// ============================================================================
#[derive(Debug, Clone, Deserialize)]
struct OidcDiscovery {
issuer: String,
authorization_endpoint: String,
token_endpoint: String,
userinfo_endpoint: Option<String>,
jwks_uri: String,
}
// ============================================================================
// JWKS structures for RS256 validation
// ============================================================================
#[derive(Debug, Clone, Deserialize)]
struct JwksDocument {
keys: Vec<JwkKey>,
}
#[derive(Debug, Clone, Deserialize)]
struct JwkKey {
kty: String,
#[serde(rename = "use")]
key_use: Option<String>,
kid: Option<String>,
alg: Option<String>,
n: Option<String>, // RSA modulus (base64url)
e: Option<String>, // RSA exponent (base64url)
}
// ============================================================================
// Token exchange response
// ============================================================================
#[derive(Debug, Deserialize)]
struct TokenResponse {
access_token: String,
id_token: Option<String>,
refresh_token: Option<String>,
#[allow(dead_code)]
token_type: Option<String>,
#[allow(dead_code)]
expires_in: Option<i64>,
}
// ============================================================================
// ID token claims (standard OIDC)
// ============================================================================
#[derive(Debug, Deserialize)]
struct IdTokenClaims {
sub: String,
email: Option<String>,
preferred_username: Option<String>,
name: Option<String>,
groups: Option<Vec<String>>,
nonce: Option<String>,
// Standard JWT fields
#[allow(dead_code)]
iss: Option<String>,
#[allow(dead_code)]
aud: Option<serde_json::Value>,
#[allow(dead_code)]
exp: Option<i64>,
#[allow(dead_code)]
iat: Option<i64>,
}
// ============================================================================
// UserInfo response
// ============================================================================
#[derive(Debug, Deserialize)]
struct UserInfoResponse {
sub: String,
email: Option<String>,
preferred_username: Option<String>,
name: Option<String>,
groups: Option<Vec<String>>,
}
// ============================================================================
// OIDC Service
// ============================================================================
pub struct OidcService {
config: OidcConfig,
http_client: reqwest::Client,
/// Cached discovery document
discovery: RwLock<Option<OidcDiscovery>>,
/// Cached JWKS
jwks: RwLock<Option<JwksDocument>>,
}
impl OidcService {
pub fn new(config: OidcConfig) -> Self {
let http_client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.expect("Failed to build HTTP client for OIDC");
Self {
config,
http_client,
discovery: RwLock::new(None),
jwks: RwLock::new(None),
}
}
/// Fetch and cache the OIDC discovery document
async fn get_discovery(&self) -> Result<OidcDiscovery, DomainError> {
// Check cache first
{
let cache = self
.discovery
.read()
.map_err(|_| DomainError::new(ErrorKind::InternalError, "OIDC", "Lock poisoned"))?;
if let Some(ref disc) = *cache {
return Ok(disc.clone());
}
}
// Fetch discovery document
let issuer = self.config.issuer_url.trim_end_matches('/');
let discovery_url = format!("{}/.well-known/openid-configuration", issuer);
tracing::info!("Fetching OIDC discovery from: {}", discovery_url);
let resp = self
.http_client
.get(&discovery_url)
.send()
.await
.map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
format!("Failed to fetch OIDC discovery: {}", e),
)
})?;
if !resp.status().is_success() {
return Err(DomainError::new(
ErrorKind::InternalError,
"OIDC",
format!("OIDC discovery returned status {}", resp.status()),
));
}
let discovery: OidcDiscovery = resp.json().await.map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
format!("Failed to parse OIDC discovery: {}", e),
)
})?;
// Cache it
{
let mut cache = self
.discovery
.write()
.map_err(|_| DomainError::new(ErrorKind::InternalError, "OIDC", "Lock poisoned"))?;
*cache = Some(discovery.clone());
}
Ok(discovery)
}
/// Fetch and cache JWKS document for ID token validation
async fn get_jwks(&self) -> Result<JwksDocument, DomainError> {
// Check cache first
{
let cache = self
.jwks
.read()
.map_err(|_| DomainError::new(ErrorKind::InternalError, "OIDC", "Lock poisoned"))?;
if let Some(ref jwks) = *cache {
return Ok(jwks.clone());
}
}
let discovery = self.get_discovery().await?;
tracing::debug!("Fetching JWKS from: {}", discovery.jwks_uri);
let resp = self
.http_client
.get(&discovery.jwks_uri)
.send()
.await
.map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
format!("Failed to fetch JWKS: {}", e),
)
})?;
let jwks: JwksDocument = resp.json().await.map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
format!("Failed to parse JWKS: {}", e),
)
})?;
// Cache it
{
let mut cache = self
.jwks
.write()
.map_err(|_| DomainError::new(ErrorKind::InternalError, "OIDC", "Lock poisoned"))?;
*cache = Some(jwks.clone());
}
Ok(jwks)
}
/// Find the right RSA key from JWKS by kid header
fn find_rsa_key<'a>(jwks: &'a JwksDocument, kid: Option<&str>) -> Option<&'a JwkKey> {
jwks.keys.iter().find(|k| {
k.kty == "RSA"
&& k.key_use.as_deref() != Some("enc") // exclude encryption keys
&& (kid.is_none() || k.kid.as_deref() == kid)
})
}
/// Extract the `kid` from a JWT header without full validation
fn extract_jwt_kid(token: &str) -> Option<String> {
let parts: Vec<&str> = token.splitn(3, '.').collect();
if parts.len() < 2 {
return None;
}
use base64::Engine;
let engine = base64::engine::general_purpose::URL_SAFE_NO_PAD;
let header_bytes = engine.decode(parts[0]).ok()?;
let header: serde_json::Value = serde_json::from_slice(&header_bytes).ok()?;
header
.get("kid")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
}
}
#[async_trait]
impl OidcServicePort for OidcService {
async fn get_authorize_url(
&self,
state: &str,
nonce: &str,
pkce_challenge: &str,
) -> Result<String, DomainError> {
// Fetch or use cached discovery to get the correct authorization_endpoint
let discovery = self.get_discovery().await?;
let auth_endpoint = discovery.authorization_endpoint;
let scopes = self.config.scopes.replace(',', " ");
let url = format!(
"{}?response_type=code&client_id={}&redirect_uri={}&scope={}&state={}&nonce={}&code_challenge={}&code_challenge_method=S256",
auth_endpoint,
urlencoding::encode(&self.config.client_id),
urlencoding::encode(&self.config.redirect_uri),
urlencoding::encode(&scopes),
urlencoding::encode(state),
urlencoding::encode(nonce),
urlencoding::encode(pkce_challenge),
);
Ok(url)
}
async fn exchange_code(
&self,
code: &str,
pkce_verifier: &str,
) -> Result<OidcTokenSet, DomainError> {
let discovery = self.get_discovery().await?;
tracing::debug!(
"Exchanging authorization code at: {}",
discovery.token_endpoint
);
let resp = self
.http_client
.post(&discovery.token_endpoint)
.form(&[
("grant_type", "authorization_code"),
("code", code),
("redirect_uri", &self.config.redirect_uri),
("client_id", &self.config.client_id),
("client_secret", &self.config.client_secret),
("code_verifier", pkce_verifier),
])
.send()
.await
.map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
format!("Token exchange failed: {}", e),
)
})?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
tracing::error!(
"OIDC token exchange error: status={}, body={}",
status,
body
);
return Err(DomainError::new(
ErrorKind::AccessDenied,
"OIDC",
format!("Token exchange failed with status {}", status),
));
}
let token_resp: TokenResponse = resp.json().await.map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
format!("Failed to parse token response: {}", e),
)
})?;
let id_token = token_resp.id_token.ok_or_else(|| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
"No id_token in token response",
)
})?;
Ok(OidcTokenSet {
access_token: token_resp.access_token,
id_token,
refresh_token: token_resp.refresh_token,
})
}
async fn validate_id_token(
&self,
id_token: &str,
expected_nonce: Option<&str>,
) -> Result<OidcIdClaims, DomainError> {
let jwks = self.get_jwks().await?;
let discovery = self.get_discovery().await?;
// Extract kid from JWT header
let kid = Self::extract_jwt_kid(id_token);
// Find the matching RSA key
let jwk = Self::find_rsa_key(&jwks, kid.as_deref()).ok_or_else(|| {
DomainError::new(
ErrorKind::AccessDenied,
"OIDC",
"No suitable RSA key found in JWKS for ID token validation",
)
})?;
let n = jwk.n.as_ref().ok_or_else(|| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
"JWKS key missing 'n' component",
)
})?;
let e = jwk.e.as_ref().ok_or_else(|| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
"JWKS key missing 'e' component",
)
})?;
// Build decoding key from RSA components
let decoding_key = jsonwebtoken::DecodingKey::from_rsa_components(n, e).map_err(|err| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
format!("Failed to build RSA decoding key: {}", err),
)
})?;
// Determine algorithm from JWKS (default RS256)
let alg = match jwk.alg.as_deref() {
Some("RS384") => jsonwebtoken::Algorithm::RS384,
Some("RS512") => jsonwebtoken::Algorithm::RS512,
_ => jsonwebtoken::Algorithm::RS256,
};
// Build validation: check expiry and issuer
let mut validation = jsonwebtoken::Validation::new(alg);
validation.set_issuer(&[&discovery.issuer]);
validation.set_audience(&[&self.config.client_id]);
let token_data =
jsonwebtoken::decode::<IdTokenClaims>(id_token, &decoding_key, &validation).map_err(
|e| {
tracing::warn!("OIDC ID token validation failed: {}", e);
DomainError::new(
ErrorKind::AccessDenied,
"OIDC",
format!("ID token validation failed: {}", e),
)
},
)?;
let claims = token_data.claims;
// Verify nonce to prevent token replay attacks
if let Some(expected) = expected_nonce {
match &claims.nonce {
Some(actual) if actual == expected => { /* OK */ }
Some(actual) => {
tracing::warn!("OIDC nonce mismatch: expected={}, got={}", expected, actual);
return Err(DomainError::new(
ErrorKind::AccessDenied,
"OIDC",
"ID token nonce mismatch — possible replay attack",
));
}
None => {
tracing::warn!("OIDC nonce missing from ID token (expected={})", expected);
// Some providers don't include nonce; log warning but don't fail
}
}
}
Ok(OidcIdClaims {
sub: claims.sub,
email: claims.email,
preferred_username: claims.preferred_username,
name: claims.name,
groups: claims.groups.unwrap_or_default(),
})
}
async fn fetch_user_info(&self, access_token: &str) -> Result<OidcIdClaims, DomainError> {
let discovery = self.get_discovery().await?;
let userinfo_url = discovery.userinfo_endpoint.ok_or_else(|| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
"No userinfo_endpoint in OIDC discovery",
)
})?;
let resp = self
.http_client
.get(&userinfo_url)
.header("Authorization", format!("Bearer {}", access_token))
.send()
.await
.map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
format!("UserInfo request failed: {}", e),
)
})?;
if !resp.status().is_success() {
return Err(DomainError::new(
ErrorKind::AccessDenied,
"OIDC",
format!("UserInfo returned status {}", resp.status()),
));
}
let info: UserInfoResponse = resp.json().await.map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"OIDC",
format!("Failed to parse UserInfo: {}", e),
)
})?;
Ok(OidcIdClaims {
sub: info.sub,
email: info.email,
preferred_username: info.preferred_username,
name: info.name,
groups: info.groups.unwrap_or_default(),
})
}
fn provider_name(&self) -> &str {
&self.config.provider_name
}
}
// We need urlencoding — let's use a minimal inline implementation
mod urlencoding {
pub fn encode(input: &str) -> String {
let mut result = String::with_capacity(input.len() * 3);
for byte in input.bytes() {
match byte {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
result.push(byte as char);
}
_ => {
result.push('%');
result.push_str(&format!("{:02X}", byte));
}
}
}
result
}
}
+115 -91
View File
@@ -1,91 +1,115 @@
//! Argon2-based password hasher implementation.
//!
//! This module provides a secure password hashing implementation using the Argon2id
//! algorithm, which is the recommended choice for password hashing as of 2023+.
use argon2::{Argon2, PasswordHash, PasswordHasher, PasswordVerifier};
use argon2::password_hash::SaltString;
use rand_core::OsRng;
use crate::application::ports::auth_ports::PasswordHasherPort;
use crate::common::errors::{DomainError, ErrorKind};
/// Argon2-based implementation of the PasswordHasherPort.
///
/// Uses Argon2id algorithm which provides resistance against both side-channel
/// and GPU-based attacks. This is the recommended algorithm for password hashing.
#[derive(Debug, Clone)]
pub struct Argon2PasswordHasher {
/// Argon2 hasher instance - uses default secure parameters
_private: (),
}
impl Argon2PasswordHasher {
/// Create a new Argon2PasswordHasher with default secure parameters.
pub fn new() -> Self {
Self { _private: () }
}
}
impl Default for Argon2PasswordHasher {
fn default() -> Self {
Self::new()
}
}
impl PasswordHasherPort for Argon2PasswordHasher {
fn hash_password(&self, password: &str) -> Result<String, DomainError> {
let salt = SaltString::generate(&mut OsRng);
let argon2 = Argon2::default();
argon2.hash_password(password.as_bytes(), &salt)
.map(|hash| hash.to_string())
.map_err(|e| DomainError::new(
ErrorKind::InternalError,
"PasswordHasher",
format!("Error generating password hash: {}", e)
))
}
fn verify_password(&self, password: &str, hash: &str) -> Result<bool, DomainError> {
let parsed_hash = PasswordHash::new(hash)
.map_err(|e| DomainError::new(
ErrorKind::InternalError,
"PasswordHasher",
format!("Error processing password hash: {}", e)
))?;
Ok(Argon2::default().verify_password(password.as_bytes(), &parsed_hash).is_ok())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_hash_and_verify_password() {
let hasher = Argon2PasswordHasher::new();
let password = "test_password_123";
let hash = hasher.hash_password(password).expect("Should hash password");
assert!(hasher.verify_password(password, &hash).expect("Should verify"));
assert!(!hasher.verify_password("wrong_password", &hash).expect("Should verify"));
}
#[test]
fn test_different_hashes_for_same_password() {
let hasher = Argon2PasswordHasher::new();
let password = "same_password";
let hash1 = hasher.hash_password(password).expect("Should hash");
let hash2 = hasher.hash_password(password).expect("Should hash");
// Hashes should be different due to random salt
assert_ne!(hash1, hash2);
// But both should verify correctly
assert!(hasher.verify_password(password, &hash1).expect("Should verify"));
assert!(hasher.verify_password(password, &hash2).expect("Should verify"));
}
}
//! Argon2-based password hasher implementation.
//!
//! This module provides a secure password hashing implementation using the Argon2id
//! algorithm, which is the recommended choice for password hashing as of 2023+.
use argon2::password_hash::SaltString;
use argon2::{Argon2, PasswordHash, PasswordHasher, PasswordVerifier};
use rand_core::OsRng;
use crate::application::ports::auth_ports::PasswordHasherPort;
use crate::common::errors::{DomainError, ErrorKind};
/// Argon2-based implementation of the PasswordHasherPort.
///
/// Uses Argon2id algorithm which provides resistance against both side-channel
/// and GPU-based attacks. This is the recommended algorithm for password hashing.
#[derive(Debug, Clone)]
pub struct Argon2PasswordHasher {
/// Argon2 hasher instance - uses default secure parameters
_private: (),
}
impl Argon2PasswordHasher {
/// Create a new Argon2PasswordHasher with default secure parameters.
pub fn new() -> Self {
Self { _private: () }
}
}
impl Default for Argon2PasswordHasher {
fn default() -> Self {
Self::new()
}
}
impl PasswordHasherPort for Argon2PasswordHasher {
fn hash_password(&self, password: &str) -> Result<String, DomainError> {
let salt = SaltString::generate(&mut OsRng);
let argon2 = Argon2::default();
argon2
.hash_password(password.as_bytes(), &salt)
.map(|hash| hash.to_string())
.map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"PasswordHasher",
format!("Error generating password hash: {}", e),
)
})
}
fn verify_password(&self, password: &str, hash: &str) -> Result<bool, DomainError> {
let parsed_hash = PasswordHash::new(hash).map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"PasswordHasher",
format!("Error processing password hash: {}", e),
)
})?;
Ok(Argon2::default()
.verify_password(password.as_bytes(), &parsed_hash)
.is_ok())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_hash_and_verify_password() {
let hasher = Argon2PasswordHasher::new();
let password = "test_password_123";
let hash = hasher
.hash_password(password)
.expect("Should hash password");
assert!(
hasher
.verify_password(password, &hash)
.expect("Should verify")
);
assert!(
!hasher
.verify_password("wrong_password", &hash)
.expect("Should verify")
);
}
#[test]
fn test_different_hashes_for_same_password() {
let hasher = Argon2PasswordHasher::new();
let password = "same_password";
let hash1 = hasher.hash_password(password).expect("Should hash");
let hash2 = hasher.hash_password(password).expect("Should hash");
// Hashes should be different due to random salt
assert_ne!(hash1, hash2);
// But both should verify correctly
assert!(
hasher
.verify_password(password, &hash1)
.expect("Should verify")
);
assert!(
hasher
.verify_password(password, &hash2)
.expect("Should verify")
);
}
}
+340 -301
View File
@@ -1,301 +1,340 @@
//! PathService - Infrastructure service for storage path management
//!
//! This service was moved from domain/services because it implements application traits
//! (StoragePort, StorageMediator) and has file system dependencies (tokio::fs).
//!
//! StoragePath (Value Object) remains in domain/services/path_service.rs
use std::path::{Path, PathBuf};
use async_trait::async_trait;
use tokio::fs;
use crate::common::errors::{DomainError, ErrorKind};
use crate::application::ports::outbound::StoragePort;
use crate::application::services::storage_mediator::{StorageMediator, StorageMediatorResult, StorageMediatorError};
use crate::domain::entities::folder::Folder;
use crate::domain::services::path_service::StoragePath;
/// Infrastructure service for handling storage path operations
pub struct PathService {
root_path: PathBuf,
}
impl PathService {
/// Creates a new path service with a specific root
pub fn new(root_path: PathBuf) -> Self {
Self { root_path }
}
/// Converts a domain path to an absolute physical path
pub fn resolve_path(&self, storage_path: &StoragePath) -> PathBuf {
let mut path = self.root_path.clone();
for segment in storage_path.segments() {
path.push(segment);
}
path
}
/// Converts a physical path to a domain path
pub fn to_storage_path(&self, physical_path: &Path) -> Option<StoragePath> {
physical_path.strip_prefix(&self.root_path).ok().map(|rel_path| {
let segments: Vec<String> = rel_path
.components()
.filter_map(|c| match c {
std::path::Component::Normal(os_str) => Some(os_str.to_string_lossy().to_string()),
_ => None,
})
.collect();
StoragePath::new(segments)
})
}
/// Creates a file path within a folder
pub fn create_file_path(&self, folder_path: &StoragePath, file_name: &str) -> StoragePath {
folder_path.join(file_name)
}
/// Checks if a path is a direct child of another
pub fn is_direct_child(&self, parent_path: &StoragePath, potential_child: &StoragePath) -> bool {
if let Some(child_parent) = potential_child.parent() {
&child_parent == parent_path
} else {
parent_path.is_empty()
}
}
/// Checks if a path is at the root
pub fn is_in_root(&self, path: &StoragePath) -> bool {
path.parent().is_none_or(|p| p.is_empty())
}
/// Gets the root path used by this service
pub fn get_root_path(&self) -> &Path {
&self.root_path
}
/// Validates a path to ensure it doesn't contain dangerous components
pub fn validate_path(&self, path: &StoragePath) -> Result<(), DomainError> {
// Check for empty segments
if path.segments().iter().any(|s| s.is_empty()) {
return Err(DomainError::new(
ErrorKind::InvalidInput,
"Path",
format!("Path contains empty segments: {}", path.to_string())
));
}
// Check for dangerous characters
let dangerous_chars = ['\\', ':', '*', '?', '"', '<', '>', '|'];
for segment in path.segments() {
if segment.contains(&dangerous_chars[..]) {
return Err(DomainError::new(
ErrorKind::InvalidInput,
"Path",
format!("Path contains dangerous characters: {}", segment)
));
}
// Check that it doesn't start with . (hidden in Unix)
if segment.starts_with('.') && segment != ".well-known" {
return Err(DomainError::new(
ErrorKind::InvalidInput,
"Path",
format!("Path segments cannot start with dot: {}", segment)
));
}
}
Ok(())
}
}
#[async_trait]
impl StoragePort for PathService {
fn resolve_path(&self, storage_path: &StoragePath) -> PathBuf {
let mut path = self.root_path.clone();
for segment in storage_path.segments() {
path.push(segment);
}
path
}
async fn ensure_directory(&self, storage_path: &StoragePath) -> Result<(), DomainError> {
// First validate the path
self.validate_path(storage_path)?;
// Resolve to physical path
let physical_path = self.resolve_path(storage_path);
// Create directories if they don't exist
if !physical_path.exists() {
fs::create_dir_all(&physical_path).await
.map_err(|e| DomainError::new(
ErrorKind::AccessDenied,
"Storage",
format!("Failed to create directory: {}", physical_path.display())
).with_source(e))?;
tracing::debug!("Created directory: {}", physical_path.display());
} else if !physical_path.is_dir() {
return Err(DomainError::new(
ErrorKind::InvalidInput,
"Storage",
format!("Path exists but is not a directory: {}", physical_path.display())
));
}
Ok(())
}
async fn file_exists(&self, storage_path: &StoragePath) -> Result<bool, DomainError> {
let physical_path = self.resolve_path(storage_path);
let exists = physical_path.exists() && physical_path.is_file();
Ok(exists)
}
async fn directory_exists(&self, storage_path: &StoragePath) -> Result<bool, DomainError> {
let physical_path = self.resolve_path(storage_path);
let exists = physical_path.exists() && physical_path.is_dir();
Ok(exists)
}
}
#[async_trait]
impl StorageMediator for PathService {
async fn get_folder_path(&self, folder_id: &str) -> StorageMediatorResult<PathBuf> {
// This is a simplified implementation since PathService doesn't have direct
// access to folder repository. It's typically used through a proxy.
Err(StorageMediatorError::NotFound(format!("Folder with ID {} not found", folder_id)))
}
async fn get_folder_storage_path(&self, folder_id: &str) -> StorageMediatorResult<StoragePath> {
// Simplified implementation - should be overridden by actual implementations
Err(StorageMediatorError::NotFound(format!("Folder with ID {} not found", folder_id)))
}
async fn get_folder(&self, folder_id: &str) -> StorageMediatorResult<Folder> {
// Simplified implementation - should be overridden by actual implementations
Err(StorageMediatorError::NotFound(format!("Folder with ID {} not found", folder_id)))
}
async fn file_exists_at_path(&self, path: &Path) -> StorageMediatorResult<bool> {
let abs_path = self.resolve_path(&StoragePath::from_string(&path.to_string_lossy()));
Ok(abs_path.exists() && abs_path.is_file())
}
async fn file_exists_at_storage_path(&self, storage_path: &StoragePath) -> StorageMediatorResult<bool> {
let abs_path = self.resolve_path(storage_path);
Ok(abs_path.exists() && abs_path.is_file())
}
async fn folder_exists_at_path(&self, path: &Path) -> StorageMediatorResult<bool> {
let abs_path = self.resolve_path(&StoragePath::from_string(&path.to_string_lossy()));
Ok(abs_path.exists() && abs_path.is_dir())
}
async fn folder_exists_at_storage_path(&self, storage_path: &StoragePath) -> StorageMediatorResult<bool> {
let abs_path = self.resolve_path(storage_path);
Ok(abs_path.exists() && abs_path.is_dir())
}
fn resolve_path(&self, relative_path: &Path) -> PathBuf {
// Convert path to storage path then resolve
let path_str = relative_path.to_string_lossy().to_string();
let storage_path = StoragePath::from_string(&path_str);
PathService::resolve_path(self, &storage_path)
}
fn resolve_storage_path(&self, storage_path: &StoragePath) -> PathBuf {
PathService::resolve_path(self, storage_path)
}
async fn ensure_directory(&self, path: &Path) -> StorageMediatorResult<()> {
let abs_path = PathService::resolve_path(self, &StoragePath::from_string(&path.to_string_lossy()));
if !abs_path.exists() {
fs::create_dir_all(&abs_path).await
.map_err(|e| StorageMediatorError::AccessError(format!("Failed to create directory: {}", e)))?;
} else if !abs_path.is_dir() {
return Err(StorageMediatorError::InvalidPath(
format!("Path exists but is not a directory: {}", abs_path.display())
));
}
Ok(())
}
async fn ensure_storage_directory(&self, storage_path: &StoragePath) -> StorageMediatorResult<()> {
let abs_path = PathService::resolve_path(self, storage_path);
if !abs_path.exists() {
fs::create_dir_all(&abs_path).await
.map_err(|e| StorageMediatorError::AccessError(format!("Failed to create directory: {}", e)))?;
} else if !abs_path.is_dir() {
return Err(StorageMediatorError::InvalidPath(
format!("Path exists but is not a directory: {}", abs_path.display())
));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_resolve_path() {
let service = PathService::new(PathBuf::from("/storage"));
let storage_path = StoragePath::from_string("test/file.txt");
let absolute = service.resolve_path(&storage_path);
assert_eq!(absolute, PathBuf::from("/storage/test/file.txt"));
}
#[test]
fn test_to_storage_path() {
let service = PathService::new(PathBuf::from("/storage"));
let physical_path = PathBuf::from("/storage/folder/file.txt");
let storage_path = service.to_storage_path(&physical_path).unwrap();
assert_eq!(storage_path.to_string(), "/folder/file.txt");
}
#[test]
fn test_is_in_root() {
let service = PathService::new(PathBuf::from("/storage"));
let root_path = StoragePath::from_string("file.txt");
let nested_path = StoragePath::from_string("folder/file.txt");
assert!(service.is_in_root(&root_path));
assert!(!service.is_in_root(&nested_path));
}
#[test]
fn test_is_direct_child() {
let service = PathService::new(PathBuf::from("/storage"));
let parent = StoragePath::from_string("folder");
let child = StoragePath::from_string("folder/file.txt");
let not_child = StoragePath::from_string("folder2/file.txt");
assert!(service.is_direct_child(&parent, &child));
assert!(!service.is_direct_child(&parent, &not_child));
}
#[test]
fn test_create_file_path() {
let service = PathService::new(PathBuf::from("/storage"));
let folder_path = StoragePath::from_string("folder");
let file_path = service.create_file_path(&folder_path, "file.txt");
assert_eq!(file_path.to_string(), "/folder/file.txt");
}
}
//! PathService - Infrastructure service for storage path management
//!
//! This service was moved from domain/services because it implements application traits
//! (StoragePort, StorageMediator) and has file system dependencies (tokio::fs).
//!
//! StoragePath (Value Object) remains in domain/services/path_service.rs
use async_trait::async_trait;
use std::path::{Path, PathBuf};
use tokio::fs;
use crate::application::ports::outbound::StoragePort;
use crate::application::services::storage_mediator::{
StorageMediator, StorageMediatorError, StorageMediatorResult,
};
use crate::common::errors::{DomainError, ErrorKind};
use crate::domain::entities::folder::Folder;
use crate::domain::services::path_service::StoragePath;
/// Infrastructure service for handling storage path operations
pub struct PathService {
root_path: PathBuf,
}
impl PathService {
/// Creates a new path service with a specific root
pub fn new(root_path: PathBuf) -> Self {
Self { root_path }
}
/// Converts a domain path to an absolute physical path
pub fn resolve_path(&self, storage_path: &StoragePath) -> PathBuf {
let mut path = self.root_path.clone();
for segment in storage_path.segments() {
path.push(segment);
}
path
}
/// Converts a physical path to a domain path
pub fn to_storage_path(&self, physical_path: &Path) -> Option<StoragePath> {
physical_path
.strip_prefix(&self.root_path)
.ok()
.map(|rel_path| {
let segments: Vec<String> = rel_path
.components()
.filter_map(|c| match c {
std::path::Component::Normal(os_str) => {
Some(os_str.to_string_lossy().to_string())
}
_ => None,
})
.collect();
StoragePath::new(segments)
})
}
/// Creates a file path within a folder
pub fn create_file_path(&self, folder_path: &StoragePath, file_name: &str) -> StoragePath {
folder_path.join(file_name)
}
/// Checks if a path is a direct child of another
pub fn is_direct_child(
&self,
parent_path: &StoragePath,
potential_child: &StoragePath,
) -> bool {
if let Some(child_parent) = potential_child.parent() {
&child_parent == parent_path
} else {
parent_path.is_empty()
}
}
/// Checks if a path is at the root
pub fn is_in_root(&self, path: &StoragePath) -> bool {
path.parent().is_none_or(|p| p.is_empty())
}
/// Gets the root path used by this service
pub fn get_root_path(&self) -> &Path {
&self.root_path
}
/// Validates a path to ensure it doesn't contain dangerous components
pub fn validate_path(&self, path: &StoragePath) -> Result<(), DomainError> {
// Check for empty segments
if path.segments().iter().any(|s| s.is_empty()) {
return Err(DomainError::new(
ErrorKind::InvalidInput,
"Path",
format!("Path contains empty segments: {}", path.to_string()),
));
}
// Check for dangerous characters
let dangerous_chars = ['\\', ':', '*', '?', '"', '<', '>', '|'];
for segment in path.segments() {
if segment.contains(&dangerous_chars[..]) {
return Err(DomainError::new(
ErrorKind::InvalidInput,
"Path",
format!("Path contains dangerous characters: {}", segment),
));
}
// Check that it doesn't start with . (hidden in Unix)
if segment.starts_with('.') && segment != ".well-known" {
return Err(DomainError::new(
ErrorKind::InvalidInput,
"Path",
format!("Path segments cannot start with dot: {}", segment),
));
}
}
Ok(())
}
}
#[async_trait]
impl StoragePort for PathService {
fn resolve_path(&self, storage_path: &StoragePath) -> PathBuf {
let mut path = self.root_path.clone();
for segment in storage_path.segments() {
path.push(segment);
}
path
}
async fn ensure_directory(&self, storage_path: &StoragePath) -> Result<(), DomainError> {
// First validate the path
self.validate_path(storage_path)?;
// Resolve to physical path
let physical_path = self.resolve_path(storage_path);
// Create directories if they don't exist
if !physical_path.exists() {
fs::create_dir_all(&physical_path).await.map_err(|e| {
DomainError::new(
ErrorKind::AccessDenied,
"Storage",
format!("Failed to create directory: {}", physical_path.display()),
)
.with_source(e)
})?;
tracing::debug!("Created directory: {}", physical_path.display());
} else if !physical_path.is_dir() {
return Err(DomainError::new(
ErrorKind::InvalidInput,
"Storage",
format!(
"Path exists but is not a directory: {}",
physical_path.display()
),
));
}
Ok(())
}
async fn file_exists(&self, storage_path: &StoragePath) -> Result<bool, DomainError> {
let physical_path = self.resolve_path(storage_path);
let exists = physical_path.exists() && physical_path.is_file();
Ok(exists)
}
async fn directory_exists(&self, storage_path: &StoragePath) -> Result<bool, DomainError> {
let physical_path = self.resolve_path(storage_path);
let exists = physical_path.exists() && physical_path.is_dir();
Ok(exists)
}
}
#[async_trait]
impl StorageMediator for PathService {
async fn get_folder_path(&self, folder_id: &str) -> StorageMediatorResult<PathBuf> {
// This is a simplified implementation since PathService doesn't have direct
// access to folder repository. It's typically used through a proxy.
Err(StorageMediatorError::NotFound(format!(
"Folder with ID {} not found",
folder_id
)))
}
async fn get_folder_storage_path(&self, folder_id: &str) -> StorageMediatorResult<StoragePath> {
// Simplified implementation - should be overridden by actual implementations
Err(StorageMediatorError::NotFound(format!(
"Folder with ID {} not found",
folder_id
)))
}
async fn get_folder(&self, folder_id: &str) -> StorageMediatorResult<Folder> {
// Simplified implementation - should be overridden by actual implementations
Err(StorageMediatorError::NotFound(format!(
"Folder with ID {} not found",
folder_id
)))
}
async fn file_exists_at_path(&self, path: &Path) -> StorageMediatorResult<bool> {
let abs_path = self.resolve_path(&StoragePath::from_string(&path.to_string_lossy()));
Ok(abs_path.exists() && abs_path.is_file())
}
async fn file_exists_at_storage_path(
&self,
storage_path: &StoragePath,
) -> StorageMediatorResult<bool> {
let abs_path = self.resolve_path(storage_path);
Ok(abs_path.exists() && abs_path.is_file())
}
async fn folder_exists_at_path(&self, path: &Path) -> StorageMediatorResult<bool> {
let abs_path = self.resolve_path(&StoragePath::from_string(&path.to_string_lossy()));
Ok(abs_path.exists() && abs_path.is_dir())
}
async fn folder_exists_at_storage_path(
&self,
storage_path: &StoragePath,
) -> StorageMediatorResult<bool> {
let abs_path = self.resolve_path(storage_path);
Ok(abs_path.exists() && abs_path.is_dir())
}
fn resolve_path(&self, relative_path: &Path) -> PathBuf {
// Convert path to storage path then resolve
let path_str = relative_path.to_string_lossy().to_string();
let storage_path = StoragePath::from_string(&path_str);
PathService::resolve_path(self, &storage_path)
}
fn resolve_storage_path(&self, storage_path: &StoragePath) -> PathBuf {
PathService::resolve_path(self, storage_path)
}
async fn ensure_directory(&self, path: &Path) -> StorageMediatorResult<()> {
let abs_path =
PathService::resolve_path(self, &StoragePath::from_string(&path.to_string_lossy()));
if !abs_path.exists() {
fs::create_dir_all(&abs_path).await.map_err(|e| {
StorageMediatorError::AccessError(format!("Failed to create directory: {}", e))
})?;
} else if !abs_path.is_dir() {
return Err(StorageMediatorError::InvalidPath(format!(
"Path exists but is not a directory: {}",
abs_path.display()
)));
}
Ok(())
}
async fn ensure_storage_directory(
&self,
storage_path: &StoragePath,
) -> StorageMediatorResult<()> {
let abs_path = PathService::resolve_path(self, storage_path);
if !abs_path.exists() {
fs::create_dir_all(&abs_path).await.map_err(|e| {
StorageMediatorError::AccessError(format!("Failed to create directory: {}", e))
})?;
} else if !abs_path.is_dir() {
return Err(StorageMediatorError::InvalidPath(format!(
"Path exists but is not a directory: {}",
abs_path.display()
)));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_resolve_path() {
let service = PathService::new(PathBuf::from("/storage"));
let storage_path = StoragePath::from_string("test/file.txt");
let absolute = service.resolve_path(&storage_path);
assert_eq!(absolute, PathBuf::from("/storage/test/file.txt"));
}
#[test]
fn test_to_storage_path() {
let service = PathService::new(PathBuf::from("/storage"));
let physical_path = PathBuf::from("/storage/folder/file.txt");
let storage_path = service.to_storage_path(&physical_path).unwrap();
assert_eq!(storage_path.to_string(), "/folder/file.txt");
}
#[test]
fn test_is_in_root() {
let service = PathService::new(PathBuf::from("/storage"));
let root_path = StoragePath::from_string("file.txt");
let nested_path = StoragePath::from_string("folder/file.txt");
assert!(service.is_in_root(&root_path));
assert!(!service.is_in_root(&nested_path));
}
#[test]
fn test_is_direct_child() {
let service = PathService::new(PathBuf::from("/storage"));
let parent = StoragePath::from_string("folder");
let child = StoragePath::from_string("folder/file.txt");
let not_child = StoragePath::from_string("folder2/file.txt");
assert!(service.is_direct_child(&parent, &child));
assert!(!service.is_direct_child(&parent, &not_child));
}
#[test]
fn test_create_file_path() {
let service = PathService::new(PathBuf::from("/storage"));
let folder_path = StoragePath::from_string("folder");
let file_path = service.create_file_path(&folder_path, "file.txt");
assert_eq!(file_path.to_string(), "/folder/file.txt");
}
}
+420 -416
View File
@@ -1,416 +1,420 @@
/**
* Thumbnail Generation Service
*
* Generates and manages image thumbnails for fast gallery previews.
*
* Features:
* - Background thumbnail generation after upload
* - Multiple sizes (icon 150x150, preview 800x600)
* - WebP output for smaller file sizes
* - LRU cache for hot thumbnails
* - Lazy generation on first request if not pre-generated
*/
use std::path::{Path, PathBuf};
use std::sync::Arc;
use tokio::sync::RwLock;
use tokio::fs;
use image::{ImageFormat, imageops::FilterType};
use lru::LruCache;
use std::num::NonZeroUsize;
use bytes::Bytes;
use async_trait::async_trait;
use crate::application::ports::thumbnail_ports::{
ThumbnailPort,
ThumbnailSize as PortThumbnailSize,
ThumbnailStatsDto,
};
use crate::domain::errors::{DomainError, ErrorKind};
/// Thumbnail sizes supported by the system
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum ThumbnailSize {
/// Small icon for file listings (150x150)
Icon,
/// Medium preview for gallery view (400x400)
Preview,
/// Large preview for detail view (800x800)
Large,
}
impl ThumbnailSize {
/// Get the maximum dimension for this size
pub fn max_dimension(&self) -> u32 {
match self {
ThumbnailSize::Icon => 150,
ThumbnailSize::Preview => 400,
ThumbnailSize::Large => 800,
}
}
/// Get the directory name for this size
pub fn dir_name(&self) -> &'static str {
match self {
ThumbnailSize::Icon => "icon",
ThumbnailSize::Preview => "preview",
ThumbnailSize::Large => "large",
}
}
/// Get all thumbnail sizes
pub fn all() -> &'static [ThumbnailSize] {
&[ThumbnailSize::Icon, ThumbnailSize::Preview, ThumbnailSize::Large]
}
}
/// Cache key for thumbnails
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct ThumbnailCacheKey {
file_id: String,
size: ThumbnailSize,
}
/// Thumbnail service for generating and caching image thumbnails
pub struct ThumbnailService {
/// Root path for thumbnail storage
thumbnails_root: PathBuf,
/// In-memory LRU cache for hot thumbnails
cache: Arc<RwLock<LruCache<ThumbnailCacheKey, Bytes>>>,
/// Maximum cache size in bytes
max_cache_bytes: usize,
/// Current cache size in bytes
current_cache_bytes: Arc<RwLock<usize>>,
}
impl ThumbnailService {
/// Create a new thumbnail service
///
/// # Arguments
/// * `storage_root` - Root path of file storage
/// * `max_cache_entries` - Maximum number of thumbnails to cache in memory
/// * `max_cache_bytes` - Maximum total bytes to cache
pub fn new(storage_root: &Path, max_cache_entries: usize, max_cache_bytes: usize) -> Self {
let thumbnails_root = storage_root.join(".thumbnails");
Self {
thumbnails_root,
cache: Arc::new(RwLock::new(LruCache::new(
NonZeroUsize::new(max_cache_entries).unwrap_or(NonZeroUsize::new(1000).unwrap())
))),
max_cache_bytes,
current_cache_bytes: Arc::new(RwLock::new(0)),
}
}
/// Initialize the thumbnail directories
pub async fn initialize(&self) -> std::io::Result<()> {
for size in ThumbnailSize::all() {
let dir = self.thumbnails_root.join(size.dir_name());
fs::create_dir_all(&dir).await?;
}
tracing::info!("🖼️ Thumbnail service initialized at {:?}", self.thumbnails_root);
Ok(())
}
/// Check if a file is an image that can have thumbnails
pub fn is_supported_image(mime_type: &str) -> bool {
matches!(
mime_type,
"image/jpeg" | "image/jpg" | "image/png" | "image/gif" | "image/webp"
)
}
/// Get the path where a thumbnail would be stored
fn get_thumbnail_path(&self, file_id: &str, size: ThumbnailSize) -> PathBuf {
self.thumbnails_root
.join(size.dir_name())
.join(format!("{}.webp", file_id))
}
/// Check if a thumbnail exists on disk
pub async fn thumbnail_exists(&self, file_id: &str, size: ThumbnailSize) -> bool {
let path = self.get_thumbnail_path(file_id, size);
fs::metadata(&path).await.is_ok()
}
/// Get a thumbnail, generating it if needed
///
/// # Arguments
/// * `file_id` - ID of the original file
/// * `size` - Desired thumbnail size
/// * `original_path` - Path to the original image file
///
/// # Returns
/// Bytes of the thumbnail image (WebP format)
pub async fn get_thumbnail(
&self,
file_id: &str,
size: ThumbnailSize,
original_path: &Path,
) -> Result<Bytes, ThumbnailError> {
let cache_key = ThumbnailCacheKey {
file_id: file_id.to_string(),
size,
};
// Check in-memory cache first
{
let cache = self.cache.read().await;
if let Some(data) = cache.peek(&cache_key) {
tracing::debug!("🔥 Thumbnail cache HIT: {} {:?}", file_id, size);
return Ok(data.clone());
}
}
// Check if thumbnail exists on disk
let thumb_path = self.get_thumbnail_path(file_id, size);
if fs::metadata(&thumb_path).await.is_ok() {
// Load from disk
let data = fs::read(&thumb_path).await
.map_err(|e| ThumbnailError::IoError(e.to_string()))?;
let bytes = Bytes::from(data);
// Add to cache
self.add_to_cache(cache_key, bytes.clone()).await;
tracing::debug!("💾 Thumbnail loaded from disk: {} {:?}", file_id, size);
return Ok(bytes);
}
// Generate thumbnail
tracing::info!("🎨 Generating thumbnail: {} {:?}", file_id, size);
let bytes = self.generate_thumbnail(original_path, size).await?;
// Save to disk
if let Some(parent) = thumb_path.parent() {
fs::create_dir_all(parent).await
.map_err(|e| ThumbnailError::IoError(e.to_string()))?;
}
fs::write(&thumb_path, &bytes).await
.map_err(|e| ThumbnailError::IoError(e.to_string()))?;
// Add to cache
self.add_to_cache(cache_key, bytes.clone()).await;
Ok(bytes)
}
/// Generate a thumbnail from an image file
async fn generate_thumbnail(
&self,
original_path: &Path,
size: ThumbnailSize,
) -> Result<Bytes, ThumbnailError> {
let path = original_path.to_path_buf();
let max_dim = size.max_dimension();
// Run image processing in blocking thread pool
let result = tokio::task::spawn_blocking(move || -> Result<Vec<u8>, ThumbnailError> {
// Load image
let img = image::open(&path)
.map_err(|e| ThumbnailError::ImageError(e.to_string()))?;
// Calculate new dimensions preserving aspect ratio
let (orig_width, orig_height) = (img.width(), img.height());
let (new_width, new_height) = if orig_width > orig_height {
let ratio = max_dim as f32 / orig_width as f32;
(max_dim, (orig_height as f32 * ratio) as u32)
} else {
let ratio = max_dim as f32 / orig_height as f32;
((orig_width as f32 * ratio) as u32, max_dim)
};
// Resize using high-quality Lanczos3 filter
let thumbnail = img.resize(new_width, new_height, FilterType::Lanczos3);
// Encode as WebP for smaller file size
let mut buffer = Vec::new();
thumbnail.write_to(
&mut std::io::Cursor::new(&mut buffer),
ImageFormat::WebP
).map_err(|e| ThumbnailError::ImageError(e.to_string()))?;
Ok(buffer)
}).await
.map_err(|e| ThumbnailError::TaskError(e.to_string()))?;
result.map(Bytes::from)
}
/// Add a thumbnail to the in-memory cache
async fn add_to_cache(&self, key: ThumbnailCacheKey, data: Bytes) {
let data_size = data.len();
// Check if adding this would exceed max cache size
let mut current_size = self.current_cache_bytes.write().await;
// Evict items if needed to make room
if *current_size + data_size > self.max_cache_bytes {
let mut cache = self.cache.write().await;
while *current_size + data_size > self.max_cache_bytes && !cache.is_empty() {
if let Some((_, evicted)) = cache.pop_lru() {
*current_size = current_size.saturating_sub(evicted.len());
}
}
}
// Add to cache
let mut cache = self.cache.write().await;
if let Some(old) = cache.put(key, data) {
*current_size = current_size.saturating_sub(old.len());
}
*current_size += data_size;
}
/// Generate all thumbnail sizes for a file in the background
///
/// This is called after file upload to pre-generate thumbnails
pub fn generate_all_sizes_background(
self: Arc<Self>,
file_id: String,
original_path: PathBuf,
) {
tokio::spawn(async move {
tracing::info!("🖼️ Background thumbnail generation starting: {}", file_id);
for size in ThumbnailSize::all() {
match self.generate_thumbnail(&original_path, *size).await {
Ok(bytes) => {
// Save to disk
let thumb_path = self.get_thumbnail_path(&file_id, *size);
if let Some(parent) = thumb_path.parent() {
let _ = fs::create_dir_all(parent).await;
}
if let Err(e) = fs::write(&thumb_path, &bytes).await {
tracing::warn!("Failed to save thumbnail {}: {}", file_id, e);
} else {
tracing::debug!("✅ Generated thumbnail: {} {:?}", file_id, size);
}
},
Err(e) => {
tracing::warn!("Failed to generate thumbnail {} {:?}: {}", file_id, size, e);
}
}
}
tracing::info!("✅ Background thumbnail generation complete: {}", file_id);
});
}
/// Delete all thumbnails for a file
pub async fn delete_thumbnails(&self, file_id: &str) -> Result<(), ThumbnailError> {
for size in ThumbnailSize::all() {
let path = self.get_thumbnail_path(file_id, *size);
if fs::metadata(&path).await.is_ok() {
fs::remove_file(&path).await
.map_err(|e| ThumbnailError::IoError(e.to_string()))?;
}
// Remove from cache
let cache_key = ThumbnailCacheKey {
file_id: file_id.to_string(),
size: *size,
};
let mut cache = self.cache.write().await;
if let Some(removed) = cache.pop(&cache_key) {
let mut current_size = self.current_cache_bytes.write().await;
*current_size = current_size.saturating_sub(removed.len());
}
}
tracing::debug!("🗑️ Deleted thumbnails for: {}", file_id);
Ok(())
}
/// Get cache statistics
pub async fn get_stats(&self) -> ThumbnailStats {
let cache = self.cache.read().await;
let current_size = *self.current_cache_bytes.read().await;
ThumbnailStats {
cached_thumbnails: cache.len(),
cache_size_bytes: current_size,
max_cache_bytes: self.max_cache_bytes,
}
}
}
// ─── Port implementation ─────────────────────────────────────────────────────
/// Convert port ThumbnailSize to infra ThumbnailSize.
impl From<PortThumbnailSize> for ThumbnailSize {
fn from(size: PortThumbnailSize) -> Self {
match size {
PortThumbnailSize::Icon => ThumbnailSize::Icon,
PortThumbnailSize::Preview => ThumbnailSize::Preview,
PortThumbnailSize::Large => ThumbnailSize::Large,
}
}
}
#[async_trait]
impl ThumbnailPort for ThumbnailService {
fn is_supported_image(&self, mime_type: &str) -> bool {
ThumbnailService::is_supported_image(mime_type)
}
async fn get_thumbnail(
&self,
file_id: &str,
size: PortThumbnailSize,
original_path: &Path,
) -> Result<Bytes, DomainError> {
self.get_thumbnail(file_id, size.into(), original_path)
.await
.map_err(|e| DomainError::new(ErrorKind::InternalError, "Thumbnail", e.to_string()))
}
fn generate_all_sizes_background(
self: Arc<Self>,
file_id: String,
original_path: PathBuf,
) {
ThumbnailService::generate_all_sizes_background(self, file_id, original_path)
}
async fn delete_thumbnails(&self, file_id: &str) -> Result<(), DomainError> {
self.delete_thumbnails(file_id)
.await
.map_err(|e| DomainError::new(ErrorKind::InternalError, "Thumbnail", e.to_string()))
}
async fn get_stats(&self) -> ThumbnailStatsDto {
let stats = self.get_stats().await;
ThumbnailStatsDto {
cached_thumbnails: stats.cached_thumbnails,
cache_size_bytes: stats.cache_size_bytes,
max_cache_bytes: stats.max_cache_bytes,
}
}
}
/// Thumbnail service errors
#[derive(Debug, thiserror::Error)]
pub enum ThumbnailError {
#[error("IO error: {0}")]
IoError(String),
#[error("Image processing error: {0}")]
ImageError(String),
#[error("Task error: {0}")]
TaskError(String),
#[error("Unsupported image format")]
UnsupportedFormat,
}
/// Statistics about the thumbnail cache
#[derive(Debug, Clone)]
pub struct ThumbnailStats {
pub cached_thumbnails: usize,
pub cache_size_bytes: usize,
pub max_cache_bytes: usize,
}
use async_trait::async_trait;
use bytes::Bytes;
use image::{ImageFormat, imageops::FilterType};
use lru::LruCache;
use std::num::NonZeroUsize;
/**
* Thumbnail Generation Service
*
* Generates and manages image thumbnails for fast gallery previews.
*
* Features:
* - Background thumbnail generation after upload
* - Multiple sizes (icon 150x150, preview 800x600)
* - WebP output for smaller file sizes
* - LRU cache for hot thumbnails
* - Lazy generation on first request if not pre-generated
*/
use std::path::{Path, PathBuf};
use std::sync::Arc;
use tokio::fs;
use tokio::sync::RwLock;
use crate::application::ports::thumbnail_ports::{
ThumbnailPort, ThumbnailSize as PortThumbnailSize, ThumbnailStatsDto,
};
use crate::domain::errors::{DomainError, ErrorKind};
/// Thumbnail sizes supported by the system
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum ThumbnailSize {
/// Small icon for file listings (150x150)
Icon,
/// Medium preview for gallery view (400x400)
Preview,
/// Large preview for detail view (800x800)
Large,
}
impl ThumbnailSize {
/// Get the maximum dimension for this size
pub fn max_dimension(&self) -> u32 {
match self {
ThumbnailSize::Icon => 150,
ThumbnailSize::Preview => 400,
ThumbnailSize::Large => 800,
}
}
/// Get the directory name for this size
pub fn dir_name(&self) -> &'static str {
match self {
ThumbnailSize::Icon => "icon",
ThumbnailSize::Preview => "preview",
ThumbnailSize::Large => "large",
}
}
/// Get all thumbnail sizes
pub fn all() -> &'static [ThumbnailSize] {
&[
ThumbnailSize::Icon,
ThumbnailSize::Preview,
ThumbnailSize::Large,
]
}
}
/// Cache key for thumbnails
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct ThumbnailCacheKey {
file_id: String,
size: ThumbnailSize,
}
/// Thumbnail service for generating and caching image thumbnails
pub struct ThumbnailService {
/// Root path for thumbnail storage
thumbnails_root: PathBuf,
/// In-memory LRU cache for hot thumbnails
cache: Arc<RwLock<LruCache<ThumbnailCacheKey, Bytes>>>,
/// Maximum cache size in bytes
max_cache_bytes: usize,
/// Current cache size in bytes
current_cache_bytes: Arc<RwLock<usize>>,
}
impl ThumbnailService {
/// Create a new thumbnail service
///
/// # Arguments
/// * `storage_root` - Root path of file storage
/// * `max_cache_entries` - Maximum number of thumbnails to cache in memory
/// * `max_cache_bytes` - Maximum total bytes to cache
pub fn new(storage_root: &Path, max_cache_entries: usize, max_cache_bytes: usize) -> Self {
let thumbnails_root = storage_root.join(".thumbnails");
Self {
thumbnails_root,
cache: Arc::new(RwLock::new(LruCache::new(
NonZeroUsize::new(max_cache_entries).unwrap_or(NonZeroUsize::new(1000).unwrap()),
))),
max_cache_bytes,
current_cache_bytes: Arc::new(RwLock::new(0)),
}
}
/// Initialize the thumbnail directories
pub async fn initialize(&self) -> std::io::Result<()> {
for size in ThumbnailSize::all() {
let dir = self.thumbnails_root.join(size.dir_name());
fs::create_dir_all(&dir).await?;
}
tracing::info!(
"🖼️ Thumbnail service initialized at {:?}",
self.thumbnails_root
);
Ok(())
}
/// Check if a file is an image that can have thumbnails
pub fn is_supported_image(mime_type: &str) -> bool {
matches!(
mime_type,
"image/jpeg" | "image/jpg" | "image/png" | "image/gif" | "image/webp"
)
}
/// Get the path where a thumbnail would be stored
fn get_thumbnail_path(&self, file_id: &str, size: ThumbnailSize) -> PathBuf {
self.thumbnails_root
.join(size.dir_name())
.join(format!("{}.webp", file_id))
}
/// Check if a thumbnail exists on disk
pub async fn thumbnail_exists(&self, file_id: &str, size: ThumbnailSize) -> bool {
let path = self.get_thumbnail_path(file_id, size);
fs::metadata(&path).await.is_ok()
}
/// Get a thumbnail, generating it if needed
///
/// # Arguments
/// * `file_id` - ID of the original file
/// * `size` - Desired thumbnail size
/// * `original_path` - Path to the original image file
///
/// # Returns
/// Bytes of the thumbnail image (WebP format)
pub async fn get_thumbnail(
&self,
file_id: &str,
size: ThumbnailSize,
original_path: &Path,
) -> Result<Bytes, ThumbnailError> {
let cache_key = ThumbnailCacheKey {
file_id: file_id.to_string(),
size,
};
// Check in-memory cache first
{
let cache = self.cache.read().await;
if let Some(data) = cache.peek(&cache_key) {
tracing::debug!("🔥 Thumbnail cache HIT: {} {:?}", file_id, size);
return Ok(data.clone());
}
}
// Check if thumbnail exists on disk
let thumb_path = self.get_thumbnail_path(file_id, size);
if fs::metadata(&thumb_path).await.is_ok() {
// Load from disk
let data = fs::read(&thumb_path)
.await
.map_err(|e| ThumbnailError::IoError(e.to_string()))?;
let bytes = Bytes::from(data);
// Add to cache
self.add_to_cache(cache_key, bytes.clone()).await;
tracing::debug!("💾 Thumbnail loaded from disk: {} {:?}", file_id, size);
return Ok(bytes);
}
// Generate thumbnail
tracing::info!("🎨 Generating thumbnail: {} {:?}", file_id, size);
let bytes = self.generate_thumbnail(original_path, size).await?;
// Save to disk
if let Some(parent) = thumb_path.parent() {
fs::create_dir_all(parent)
.await
.map_err(|e| ThumbnailError::IoError(e.to_string()))?;
}
fs::write(&thumb_path, &bytes)
.await
.map_err(|e| ThumbnailError::IoError(e.to_string()))?;
// Add to cache
self.add_to_cache(cache_key, bytes.clone()).await;
Ok(bytes)
}
/// Generate a thumbnail from an image file
async fn generate_thumbnail(
&self,
original_path: &Path,
size: ThumbnailSize,
) -> Result<Bytes, ThumbnailError> {
let path = original_path.to_path_buf();
let max_dim = size.max_dimension();
// Run image processing in blocking thread pool
let result = tokio::task::spawn_blocking(move || -> Result<Vec<u8>, ThumbnailError> {
// Load image
let img = image::open(&path).map_err(|e| ThumbnailError::ImageError(e.to_string()))?;
// Calculate new dimensions preserving aspect ratio
let (orig_width, orig_height) = (img.width(), img.height());
let (new_width, new_height) = if orig_width > orig_height {
let ratio = max_dim as f32 / orig_width as f32;
(max_dim, (orig_height as f32 * ratio) as u32)
} else {
let ratio = max_dim as f32 / orig_height as f32;
((orig_width as f32 * ratio) as u32, max_dim)
};
// Resize using high-quality Lanczos3 filter
let thumbnail = img.resize(new_width, new_height, FilterType::Lanczos3);
// Encode as WebP for smaller file size
let mut buffer = Vec::new();
thumbnail
.write_to(&mut std::io::Cursor::new(&mut buffer), ImageFormat::WebP)
.map_err(|e| ThumbnailError::ImageError(e.to_string()))?;
Ok(buffer)
})
.await
.map_err(|e| ThumbnailError::TaskError(e.to_string()))?;
result.map(Bytes::from)
}
/// Add a thumbnail to the in-memory cache
async fn add_to_cache(&self, key: ThumbnailCacheKey, data: Bytes) {
let data_size = data.len();
// Check if adding this would exceed max cache size
let mut current_size = self.current_cache_bytes.write().await;
// Evict items if needed to make room
if *current_size + data_size > self.max_cache_bytes {
let mut cache = self.cache.write().await;
while *current_size + data_size > self.max_cache_bytes && !cache.is_empty() {
if let Some((_, evicted)) = cache.pop_lru() {
*current_size = current_size.saturating_sub(evicted.len());
}
}
}
// Add to cache
let mut cache = self.cache.write().await;
if let Some(old) = cache.put(key, data) {
*current_size = current_size.saturating_sub(old.len());
}
*current_size += data_size;
}
/// Generate all thumbnail sizes for a file in the background
///
/// This is called after file upload to pre-generate thumbnails
pub fn generate_all_sizes_background(self: Arc<Self>, file_id: String, original_path: PathBuf) {
tokio::spawn(async move {
tracing::info!("🖼️ Background thumbnail generation starting: {}", file_id);
for size in ThumbnailSize::all() {
match self.generate_thumbnail(&original_path, *size).await {
Ok(bytes) => {
// Save to disk
let thumb_path = self.get_thumbnail_path(&file_id, *size);
if let Some(parent) = thumb_path.parent() {
let _ = fs::create_dir_all(parent).await;
}
if let Err(e) = fs::write(&thumb_path, &bytes).await {
tracing::warn!("Failed to save thumbnail {}: {}", file_id, e);
} else {
tracing::debug!("✅ Generated thumbnail: {} {:?}", file_id, size);
}
}
Err(e) => {
tracing::warn!(
"Failed to generate thumbnail {} {:?}: {}",
file_id,
size,
e
);
}
}
}
tracing::info!("✅ Background thumbnail generation complete: {}", file_id);
});
}
/// Delete all thumbnails for a file
pub async fn delete_thumbnails(&self, file_id: &str) -> Result<(), ThumbnailError> {
for size in ThumbnailSize::all() {
let path = self.get_thumbnail_path(file_id, *size);
if fs::metadata(&path).await.is_ok() {
fs::remove_file(&path)
.await
.map_err(|e| ThumbnailError::IoError(e.to_string()))?;
}
// Remove from cache
let cache_key = ThumbnailCacheKey {
file_id: file_id.to_string(),
size: *size,
};
let mut cache = self.cache.write().await;
if let Some(removed) = cache.pop(&cache_key) {
let mut current_size = self.current_cache_bytes.write().await;
*current_size = current_size.saturating_sub(removed.len());
}
}
tracing::debug!("🗑️ Deleted thumbnails for: {}", file_id);
Ok(())
}
/// Get cache statistics
pub async fn get_stats(&self) -> ThumbnailStats {
let cache = self.cache.read().await;
let current_size = *self.current_cache_bytes.read().await;
ThumbnailStats {
cached_thumbnails: cache.len(),
cache_size_bytes: current_size,
max_cache_bytes: self.max_cache_bytes,
}
}
}
// ─── Port implementation ─────────────────────────────────────────────────────
/// Convert port ThumbnailSize to infra ThumbnailSize.
impl From<PortThumbnailSize> for ThumbnailSize {
fn from(size: PortThumbnailSize) -> Self {
match size {
PortThumbnailSize::Icon => ThumbnailSize::Icon,
PortThumbnailSize::Preview => ThumbnailSize::Preview,
PortThumbnailSize::Large => ThumbnailSize::Large,
}
}
}
#[async_trait]
impl ThumbnailPort for ThumbnailService {
fn is_supported_image(&self, mime_type: &str) -> bool {
ThumbnailService::is_supported_image(mime_type)
}
async fn get_thumbnail(
&self,
file_id: &str,
size: PortThumbnailSize,
original_path: &Path,
) -> Result<Bytes, DomainError> {
self.get_thumbnail(file_id, size.into(), original_path)
.await
.map_err(|e| DomainError::new(ErrorKind::InternalError, "Thumbnail", e.to_string()))
}
fn generate_all_sizes_background(self: Arc<Self>, file_id: String, original_path: PathBuf) {
ThumbnailService::generate_all_sizes_background(self, file_id, original_path)
}
async fn delete_thumbnails(&self, file_id: &str) -> Result<(), DomainError> {
self.delete_thumbnails(file_id)
.await
.map_err(|e| DomainError::new(ErrorKind::InternalError, "Thumbnail", e.to_string()))
}
async fn get_stats(&self) -> ThumbnailStatsDto {
let stats = self.get_stats().await;
ThumbnailStatsDto {
cached_thumbnails: stats.cached_thumbnails,
cache_size_bytes: stats.cache_size_bytes,
max_cache_bytes: stats.max_cache_bytes,
}
}
}
/// Thumbnail service errors
#[derive(Debug, thiserror::Error)]
pub enum ThumbnailError {
#[error("IO error: {0}")]
IoError(String),
#[error("Image processing error: {0}")]
ImageError(String),
#[error("Task error: {0}")]
TaskError(String),
#[error("Unsupported image format")]
UnsupportedFormat,
}
/// Statistics about the thumbnail cache
#[derive(Debug, Clone)]
pub struct ThumbnailStats {
pub cached_thumbnails: usize,
pub cache_size_bytes: usize,
pub max_cache_bytes: usize,
}
@@ -3,9 +3,9 @@ use std::time::Duration;
use tokio::time;
use tracing::{debug, error, info, instrument};
use crate::application::ports::trash_ports::TrashUseCase;
use crate::common::errors::Result;
use crate::domain::repositories::trash_repository::TrashRepository;
use crate::application::ports::trash_ports::TrashUseCase;
/// Service for automatic cleanup of expired items in the trash
pub struct TrashCleanupService {
@@ -26,38 +26,42 @@ impl TrashCleanupService {
cleanup_interval_hours: cleanup_interval_hours.max(1), // Minimum 1 hour
}
}
/// Starts the periodic cleanup job
#[instrument(skip(self))]
pub async fn start_cleanup_job(&self) {
let trash_repository = self.trash_repository.clone();
let trash_service = self.trash_service.clone();
let interval_hours = self.cleanup_interval_hours;
info!("Starting trash cleanup job with interval of {} hours", interval_hours);
info!(
"Starting trash cleanup job with interval of {} hours",
interval_hours
);
tokio::spawn(async move {
let interval_duration = Duration::from_secs(interval_hours * 60 * 60);
let mut interval = time::interval(interval_duration);
// First immediate execution
Self::cleanup_expired_items(trash_repository.clone(), trash_service.clone()).await
Self::cleanup_expired_items(trash_repository.clone(), trash_service.clone())
.await
.unwrap_or_else(|e| error!("Error in initial trash cleanup: {:?}", e));
loop {
interval.tick().await;
debug!("Running scheduled trash cleanup task");
if let Err(e) = Self::cleanup_expired_items(
trash_repository.clone(),
trash_service.clone()
).await {
if let Err(e) =
Self::cleanup_expired_items(trash_repository.clone(), trash_service.clone())
.await
{
error!("Error in scheduled trash cleanup: {:?}", e);
}
}
});
}
/// Cleans up expired items in the trash
#[instrument(skip(trash_repository, trash_service))]
async fn cleanup_expired_items(
@@ -65,24 +69,24 @@ impl TrashCleanupService {
trash_service: Arc<dyn TrashUseCase>,
) -> Result<()> {
debug!("Starting cleanup of expired items in the trash");
// Get all expired items
let expired_items = trash_repository.get_expired_items().await?;
if expired_items.is_empty() {
debug!("No expired items to clean up");
return Ok(());
}
info!("Found {} expired items to delete", expired_items.len());
// Delete each expired item
for item in expired_items {
let trash_id = item.id().to_string();
let user_id = item.user_id().to_string();
debug!("Deleting expired item: id={}, user={}", trash_id, user_id);
// If a deletion fails, continue with the rest
if let Err(e) = trash_service.delete_permanently(&trash_id, &user_id).await {
error!("Error deleting expired item {}: {:?}", trash_id, e);
@@ -90,8 +94,8 @@ impl TrashCleanupService {
debug!("Expired item deleted successfully: {}", trash_id);
}
}
info!("Trash cleanup completed");
Ok(())
}
}
}
+501 -483
View File
@@ -1,483 +1,501 @@
// ═══════════════════════════════════════════════════════════════════════════════
// WRITE-BEHIND CACHE - Zero-latency uploads for small files
// ═══════════════════════════════════════════════════════════════════════════════
//
// Strategy:
// 1. For files < 1MB, store in RAM and respond immediately (201 Created)
// 2. Flush to disk asynchronously in background
// 3. Serve reads from cache while pending flush
// 4. On read miss, check if pending then serve from cache
//
// This gives users perceived ~0ms upload latency for small files
// ═══════════════════════════════════════════════════════════════════════════════
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{RwLock, mpsc};
use tokio::fs;
use tokio::io::AsyncWriteExt;
use bytes::Bytes;
use async_trait::async_trait;
use crate::application::ports::cache_ports::{WriteBehindCachePort, WriteBehindStatsDto};
use crate::domain::errors::DomainError;
/// Maximum size for write-behind cache (files larger bypass cache)
const WRITE_BEHIND_MAX_SIZE: usize = 1024 * 1024; // 1MB
/// Maximum total cache size in bytes
const MAX_CACHE_SIZE: usize = 100 * 1024 * 1024; // 100MB total
/// Maximum time a file can stay pending before forced flush
const MAX_PENDING_DURATION: Duration = Duration::from_secs(30);
/// Flush check interval
const FLUSH_INTERVAL: Duration = Duration::from_millis(100);
/// Entry in the write-behind cache
#[derive(Clone)]
pub struct PendingWrite {
/// File content
pub content: Bytes,
/// Target path on disk
pub target_path: PathBuf,
/// When this entry was created
pub created_at: Instant,
/// File ID for tracking
pub file_id: String,
}
/// Statistics for monitoring
#[derive(Debug, Clone, Default)]
pub struct WriteBehindStats {
pub pending_count: usize,
pub pending_bytes: usize,
pub total_writes: u64,
pub total_bytes_written: u64,
pub cache_hits: u64,
pub avg_flush_time_us: u64,
}
/// Write-Behind Cache for zero-latency small file uploads
pub struct WriteBehindCache {
/// Pending writes indexed by file ID
pending: Arc<RwLock<HashMap<String, PendingWrite>>>,
/// Current total size of pending data
current_size: Arc<RwLock<usize>>,
/// Channel to signal flush worker
flush_tx: mpsc::Sender<FlushCommand>,
/// Statistics
stats: Arc<RwLock<WriteBehindStats>>,
}
/// Commands for the flush worker
enum FlushCommand {
/// Flush a specific file
FlushFile(String),
/// Flush all pending files
FlushAll,
/// Shutdown the worker
Shutdown,
}
impl WriteBehindCache {
/// Create a new write-behind cache with background flush worker
pub fn new() -> Arc<Self> {
let (flush_tx, flush_rx) = mpsc::channel(1000);
let cache = Arc::new(Self {
pending: Arc::new(RwLock::new(HashMap::new())),
current_size: Arc::new(RwLock::new(0)),
flush_tx,
stats: Arc::new(RwLock::new(WriteBehindStats::default())),
});
// Start the background flush worker
let cache_clone = cache.clone();
tokio::spawn(async move {
cache_clone.flush_worker(flush_rx).await;
});
// Start the periodic flush checker
let cache_clone2 = cache.clone();
tokio::spawn(async move {
cache_clone2.periodic_flush_checker().await;
});
tracing::info!("⚡ Write-Behind Cache initialized (max {}MB)", MAX_CACHE_SIZE / (1024 * 1024));
cache
}
/// Check if a file size is eligible for write-behind caching
#[inline]
pub fn is_eligible(size: usize) -> bool {
size <= WRITE_BEHIND_MAX_SIZE
}
/// Put a file in the pending write cache
/// Returns Ok(true) if cached, Ok(false) if cache is full
pub async fn put_pending(
&self,
file_id: String,
content: Bytes,
target_path: PathBuf,
) -> Result<bool, std::io::Error> {
let content_size = content.len();
// Check if we have space
{
let current = *self.current_size.read().await;
if current + content_size > MAX_CACHE_SIZE {
tracing::debug!(
"Write-behind cache full ({}/{}MB), bypassing for {}",
current / (1024 * 1024),
MAX_CACHE_SIZE / (1024 * 1024),
file_id
);
return Ok(false);
}
}
// Add to pending
let entry = PendingWrite {
content,
target_path,
created_at: Instant::now(),
file_id: file_id.clone(),
};
{
let mut pending = self.pending.write().await;
let mut size = self.current_size.write().await;
// If replacing existing entry, adjust size
if let Some(old) = pending.insert(file_id.clone(), entry) {
*size -= old.content.len();
}
*size += content_size;
}
// Update stats
{
let mut stats = self.stats.write().await;
stats.pending_count += 1;
stats.pending_bytes += content_size;
}
// Signal flush worker (non-blocking)
let _ = self.flush_tx.try_send(FlushCommand::FlushFile(file_id.clone()));
tracing::debug!("⚡ Cached pending write: {} ({} bytes)", file_id, content_size);
Ok(true)
}
/// Get content from cache if pending (for reads before flush completes)
pub async fn get_pending(&self, file_id: &str) -> Option<Bytes> {
let pending = self.pending.read().await;
if let Some(entry) = pending.get(file_id) {
// Update cache hit stats
let mut stats = self.stats.write().await;
stats.cache_hits += 1;
tracing::debug!("⚡ Cache hit for pending file: {}", file_id);
return Some(entry.content.clone());
}
None
}
/// Check if a file is pending flush
pub async fn is_pending(&self, file_id: &str) -> bool {
self.pending.read().await.contains_key(file_id)
}
/// Force immediate flush of a specific file (for critical operations)
pub async fn force_flush(&self, file_id: &str) -> Result<(), std::io::Error> {
let entry = {
let pending = self.pending.read().await;
pending.get(file_id).cloned()
};
if let Some(entry) = entry {
self.flush_single(&entry.file_id, &entry).await?;
}
Ok(())
}
/// Flush all pending writes immediately
pub async fn flush_all(&self) -> Result<(), std::io::Error> {
let _ = self.flush_tx.send(FlushCommand::FlushAll).await;
// Wait a bit for flush to complete
tokio::time::sleep(Duration::from_millis(50)).await;
Ok(())
}
/// Gracefully shutdown the write-behind cache
/// Flushes all pending writes before stopping the background worker
pub async fn shutdown(&self) -> Result<(), std::io::Error> {
tracing::info!("🛑 Shutting down write-behind cache...");
// First flush all pending writes
self.flush_all().await?;
// Then signal the worker to stop
let _ = self.flush_tx.send(FlushCommand::Shutdown).await;
// Give worker time to process shutdown
tokio::time::sleep(Duration::from_millis(100)).await;
tracing::info!("✅ Write-behind cache shutdown complete");
Ok(())
}
/// Get current statistics
pub async fn get_stats(&self) -> WriteBehindStats {
self.stats.read().await.clone()
}
/// Background worker that handles actual disk writes
async fn flush_worker(&self, mut rx: mpsc::Receiver<FlushCommand>) {
tracing::info!("🔄 Write-behind flush worker started");
while let Some(cmd) = rx.recv().await {
match cmd {
FlushCommand::FlushFile(file_id) => {
// Small delay to batch nearby writes
tokio::time::sleep(Duration::from_millis(10)).await;
let entry = {
let pending = self.pending.read().await;
pending.get(&file_id).cloned()
};
if let Some(entry) = entry
&& let Err(e) = self.flush_single(&file_id, &entry).await {
tracing::error!("Failed to flush {}: {}", file_id, e);
// Keep in cache for retry
continue;
}
}
FlushCommand::FlushAll => {
let entries: Vec<_> = {
let pending = self.pending.read().await;
pending.iter().map(|(k, v)| (k.clone(), v.clone())).collect()
};
for (file_id, entry) in entries {
if let Err(e) = self.flush_single(&file_id, &entry).await {
tracing::error!("Failed to flush {}: {}", file_id, e);
}
}
}
FlushCommand::Shutdown => {
tracing::info!("Write-behind flush worker shutting down");
break;
}
}
}
}
/// Flush a single file to disk
async fn flush_single(&self, file_id: &str, entry: &PendingWrite) -> Result<(), std::io::Error> {
let start = Instant::now();
// Ensure parent directory exists
if let Some(parent) = entry.target_path.parent() {
fs::create_dir_all(parent).await?;
}
// Write atomically using temp file + rename
let temp_path = entry.target_path.with_extension("tmp");
{
let mut file = fs::File::create(&temp_path).await?;
file.write_all(&entry.content).await?;
file.sync_all().await?;
}
fs::rename(&temp_path, &entry.target_path).await?;
let elapsed = start.elapsed();
let content_len = entry.content.len();
// Remove from pending
{
let mut pending = self.pending.write().await;
let mut size = self.current_size.write().await;
if pending.remove(file_id).is_some() {
*size = size.saturating_sub(content_len);
}
}
// Update stats
{
let mut stats = self.stats.write().await;
stats.pending_count = stats.pending_count.saturating_sub(1);
stats.pending_bytes = stats.pending_bytes.saturating_sub(content_len);
stats.total_writes += 1;
stats.total_bytes_written += content_len as u64;
// Running average of flush time
let flush_us = elapsed.as_micros() as u64;
if stats.avg_flush_time_us == 0 {
stats.avg_flush_time_us = flush_us;
} else {
stats.avg_flush_time_us = (stats.avg_flush_time_us * 9 + flush_us) / 10;
}
}
tracing::debug!(
"💾 Flushed {} to disk ({} bytes in {:?})",
file_id,
content_len,
elapsed
);
Ok(())
}
/// Periodic checker for stale pending writes
async fn periodic_flush_checker(&self) {
let mut interval = tokio::time::interval(FLUSH_INTERVAL);
loop {
interval.tick().await;
let stale_files: Vec<String> = {
let pending = self.pending.read().await;
pending
.iter()
.filter(|(_, entry)| entry.created_at.elapsed() > MAX_PENDING_DURATION)
.map(|(id, _)| id.clone())
.collect()
};
for file_id in stale_files {
tracing::warn!("Forcing flush of stale pending file: {}", file_id);
let _ = self.flush_tx.try_send(FlushCommand::FlushFile(file_id));
}
}
}
}
// ─── Port implementation ─────────────────────────────────────────────────────
#[async_trait]
impl WriteBehindCachePort for WriteBehindCache {
fn is_eligible_size(&self, size: usize) -> bool {
WriteBehindCache::is_eligible(size)
}
async fn put_pending(
&self,
file_id: String,
content: Bytes,
target_path: PathBuf,
) -> Result<bool, DomainError> {
self.put_pending(file_id, content, target_path).await.map_err(DomainError::from)
}
async fn get_pending(&self, file_id: &str) -> Option<Bytes> {
self.get_pending(file_id).await
}
async fn is_pending(&self, file_id: &str) -> bool {
self.is_pending(file_id).await
}
async fn force_flush(&self, file_id: &str) -> Result<(), DomainError> {
self.force_flush(file_id).await.map_err(DomainError::from)
}
async fn flush_all(&self) -> Result<(), DomainError> {
self.flush_all().await.map_err(DomainError::from)
}
async fn shutdown(&self) -> Result<(), DomainError> {
self.shutdown().await.map_err(DomainError::from)
}
async fn get_stats(&self) -> WriteBehindStatsDto {
let stats = self.get_stats().await;
WriteBehindStatsDto {
pending_count: stats.pending_count,
pending_bytes: stats.pending_bytes,
total_writes: stats.total_writes,
total_bytes_written: stats.total_bytes_written,
cache_hits: stats.cache_hits,
avg_flush_time_us: stats.avg_flush_time_us,
}
}
}
impl Default for WriteBehindCache {
fn default() -> Self {
// Note: This creates a non-Arc version, prefer using new()
let (flush_tx, _) = mpsc::channel(1);
Self {
pending: Arc::new(RwLock::new(HashMap::new())),
current_size: Arc::new(RwLock::new(0)),
flush_tx,
stats: Arc::new(RwLock::new(WriteBehindStats::default())),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::TempDir;
#[tokio::test]
async fn test_write_behind_basic() {
let cache = WriteBehindCache::new();
let temp_dir = TempDir::new().unwrap();
let target = temp_dir.path().join("test.txt");
let content = Bytes::from("Hello, World!");
// Put in cache
let cached = cache.put_pending(
"test-id".to_string(),
content.clone(),
target.clone(),
).await.unwrap();
assert!(cached);
assert!(cache.is_pending("test-id").await);
// Should be readable from cache
let cached_content = cache.get_pending("test-id").await.unwrap();
assert_eq!(cached_content, content);
// Force flush
cache.force_flush("test-id").await.unwrap();
// Should no longer be pending
assert!(!cache.is_pending("test-id").await);
// File should exist on disk
assert!(target.exists());
let disk_content = std::fs::read(&target).unwrap();
assert_eq!(disk_content, content.as_ref());
}
#[tokio::test]
async fn test_eligibility() {
// 500KB should be eligible
assert!(WriteBehindCache::is_eligible(500 * 1024));
// 1MB exactly should be eligible
assert!(WriteBehindCache::is_eligible(1024 * 1024));
// Over 1MB should not be eligible
assert!(!WriteBehindCache::is_eligible(1024 * 1024 + 1));
}
}
// ═══════════════════════════════════════════════════════════════════════════════
// WRITE-BEHIND CACHE - Zero-latency uploads for small files
// ═══════════════════════════════════════════════════════════════════════════════
//
// Strategy:
// 1. For files < 1MB, store in RAM and respond immediately (201 Created)
// 2. Flush to disk asynchronously in background
// 3. Serve reads from cache while pending flush
// 4. On read miss, check if pending then serve from cache
//
// This gives users perceived ~0ms upload latency for small files
// ═══════════════════════════════════════════════════════════════════════════════
use async_trait::async_trait;
use bytes::Bytes;
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::fs;
use tokio::io::AsyncWriteExt;
use tokio::sync::{RwLock, mpsc};
use crate::application::ports::cache_ports::{WriteBehindCachePort, WriteBehindStatsDto};
use crate::domain::errors::DomainError;
/// Maximum size for write-behind cache (files larger bypass cache)
const WRITE_BEHIND_MAX_SIZE: usize = 1024 * 1024; // 1MB
/// Maximum total cache size in bytes
const MAX_CACHE_SIZE: usize = 100 * 1024 * 1024; // 100MB total
/// Maximum time a file can stay pending before forced flush
const MAX_PENDING_DURATION: Duration = Duration::from_secs(30);
/// Flush check interval
const FLUSH_INTERVAL: Duration = Duration::from_millis(100);
/// Entry in the write-behind cache
#[derive(Clone)]
pub struct PendingWrite {
/// File content
pub content: Bytes,
/// Target path on disk
pub target_path: PathBuf,
/// When this entry was created
pub created_at: Instant,
/// File ID for tracking
pub file_id: String,
}
/// Statistics for monitoring
#[derive(Debug, Clone, Default)]
pub struct WriteBehindStats {
pub pending_count: usize,
pub pending_bytes: usize,
pub total_writes: u64,
pub total_bytes_written: u64,
pub cache_hits: u64,
pub avg_flush_time_us: u64,
}
/// Write-Behind Cache for zero-latency small file uploads
pub struct WriteBehindCache {
/// Pending writes indexed by file ID
pending: Arc<RwLock<HashMap<String, PendingWrite>>>,
/// Current total size of pending data
current_size: Arc<RwLock<usize>>,
/// Channel to signal flush worker
flush_tx: mpsc::Sender<FlushCommand>,
/// Statistics
stats: Arc<RwLock<WriteBehindStats>>,
}
/// Commands for the flush worker
enum FlushCommand {
/// Flush a specific file
FlushFile(String),
/// Flush all pending files
FlushAll,
/// Shutdown the worker
Shutdown,
}
impl WriteBehindCache {
/// Create a new write-behind cache with background flush worker
pub fn new() -> Arc<Self> {
let (flush_tx, flush_rx) = mpsc::channel(1000);
let cache = Arc::new(Self {
pending: Arc::new(RwLock::new(HashMap::new())),
current_size: Arc::new(RwLock::new(0)),
flush_tx,
stats: Arc::new(RwLock::new(WriteBehindStats::default())),
});
// Start the background flush worker
let cache_clone = cache.clone();
tokio::spawn(async move {
cache_clone.flush_worker(flush_rx).await;
});
// Start the periodic flush checker
let cache_clone2 = cache.clone();
tokio::spawn(async move {
cache_clone2.periodic_flush_checker().await;
});
tracing::info!(
"⚡ Write-Behind Cache initialized (max {}MB)",
MAX_CACHE_SIZE / (1024 * 1024)
);
cache
}
/// Check if a file size is eligible for write-behind caching
#[inline]
pub fn is_eligible(size: usize) -> bool {
size <= WRITE_BEHIND_MAX_SIZE
}
/// Put a file in the pending write cache
/// Returns Ok(true) if cached, Ok(false) if cache is full
pub async fn put_pending(
&self,
file_id: String,
content: Bytes,
target_path: PathBuf,
) -> Result<bool, std::io::Error> {
let content_size = content.len();
// Check if we have space
{
let current = *self.current_size.read().await;
if current + content_size > MAX_CACHE_SIZE {
tracing::debug!(
"Write-behind cache full ({}/{}MB), bypassing for {}",
current / (1024 * 1024),
MAX_CACHE_SIZE / (1024 * 1024),
file_id
);
return Ok(false);
}
}
// Add to pending
let entry = PendingWrite {
content,
target_path,
created_at: Instant::now(),
file_id: file_id.clone(),
};
{
let mut pending = self.pending.write().await;
let mut size = self.current_size.write().await;
// If replacing existing entry, adjust size
if let Some(old) = pending.insert(file_id.clone(), entry) {
*size -= old.content.len();
}
*size += content_size;
}
// Update stats
{
let mut stats = self.stats.write().await;
stats.pending_count += 1;
stats.pending_bytes += content_size;
}
// Signal flush worker (non-blocking)
let _ = self
.flush_tx
.try_send(FlushCommand::FlushFile(file_id.clone()));
tracing::debug!(
"⚡ Cached pending write: {} ({} bytes)",
file_id,
content_size
);
Ok(true)
}
/// Get content from cache if pending (for reads before flush completes)
pub async fn get_pending(&self, file_id: &str) -> Option<Bytes> {
let pending = self.pending.read().await;
if let Some(entry) = pending.get(file_id) {
// Update cache hit stats
let mut stats = self.stats.write().await;
stats.cache_hits += 1;
tracing::debug!("⚡ Cache hit for pending file: {}", file_id);
return Some(entry.content.clone());
}
None
}
/// Check if a file is pending flush
pub async fn is_pending(&self, file_id: &str) -> bool {
self.pending.read().await.contains_key(file_id)
}
/// Force immediate flush of a specific file (for critical operations)
pub async fn force_flush(&self, file_id: &str) -> Result<(), std::io::Error> {
let entry = {
let pending = self.pending.read().await;
pending.get(file_id).cloned()
};
if let Some(entry) = entry {
self.flush_single(&entry.file_id, &entry).await?;
}
Ok(())
}
/// Flush all pending writes immediately
pub async fn flush_all(&self) -> Result<(), std::io::Error> {
let _ = self.flush_tx.send(FlushCommand::FlushAll).await;
// Wait a bit for flush to complete
tokio::time::sleep(Duration::from_millis(50)).await;
Ok(())
}
/// Gracefully shutdown the write-behind cache
/// Flushes all pending writes before stopping the background worker
pub async fn shutdown(&self) -> Result<(), std::io::Error> {
tracing::info!("🛑 Shutting down write-behind cache...");
// First flush all pending writes
self.flush_all().await?;
// Then signal the worker to stop
let _ = self.flush_tx.send(FlushCommand::Shutdown).await;
// Give worker time to process shutdown
tokio::time::sleep(Duration::from_millis(100)).await;
tracing::info!("✅ Write-behind cache shutdown complete");
Ok(())
}
/// Get current statistics
pub async fn get_stats(&self) -> WriteBehindStats {
self.stats.read().await.clone()
}
/// Background worker that handles actual disk writes
async fn flush_worker(&self, mut rx: mpsc::Receiver<FlushCommand>) {
tracing::info!("🔄 Write-behind flush worker started");
while let Some(cmd) = rx.recv().await {
match cmd {
FlushCommand::FlushFile(file_id) => {
// Small delay to batch nearby writes
tokio::time::sleep(Duration::from_millis(10)).await;
let entry = {
let pending = self.pending.read().await;
pending.get(&file_id).cloned()
};
if let Some(entry) = entry
&& let Err(e) = self.flush_single(&file_id, &entry).await
{
tracing::error!("Failed to flush {}: {}", file_id, e);
// Keep in cache for retry
continue;
}
}
FlushCommand::FlushAll => {
let entries: Vec<_> = {
let pending = self.pending.read().await;
pending
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect()
};
for (file_id, entry) in entries {
if let Err(e) = self.flush_single(&file_id, &entry).await {
tracing::error!("Failed to flush {}: {}", file_id, e);
}
}
}
FlushCommand::Shutdown => {
tracing::info!("Write-behind flush worker shutting down");
break;
}
}
}
}
/// Flush a single file to disk
async fn flush_single(
&self,
file_id: &str,
entry: &PendingWrite,
) -> Result<(), std::io::Error> {
let start = Instant::now();
// Ensure parent directory exists
if let Some(parent) = entry.target_path.parent() {
fs::create_dir_all(parent).await?;
}
// Write atomically using temp file + rename
let temp_path = entry.target_path.with_extension("tmp");
{
let mut file = fs::File::create(&temp_path).await?;
file.write_all(&entry.content).await?;
file.sync_all().await?;
}
fs::rename(&temp_path, &entry.target_path).await?;
let elapsed = start.elapsed();
let content_len = entry.content.len();
// Remove from pending
{
let mut pending = self.pending.write().await;
let mut size = self.current_size.write().await;
if pending.remove(file_id).is_some() {
*size = size.saturating_sub(content_len);
}
}
// Update stats
{
let mut stats = self.stats.write().await;
stats.pending_count = stats.pending_count.saturating_sub(1);
stats.pending_bytes = stats.pending_bytes.saturating_sub(content_len);
stats.total_writes += 1;
stats.total_bytes_written += content_len as u64;
// Running average of flush time
let flush_us = elapsed.as_micros() as u64;
if stats.avg_flush_time_us == 0 {
stats.avg_flush_time_us = flush_us;
} else {
stats.avg_flush_time_us = (stats.avg_flush_time_us * 9 + flush_us) / 10;
}
}
tracing::debug!(
"💾 Flushed {} to disk ({} bytes in {:?})",
file_id,
content_len,
elapsed
);
Ok(())
}
/// Periodic checker for stale pending writes
async fn periodic_flush_checker(&self) {
let mut interval = tokio::time::interval(FLUSH_INTERVAL);
loop {
interval.tick().await;
let stale_files: Vec<String> = {
let pending = self.pending.read().await;
pending
.iter()
.filter(|(_, entry)| entry.created_at.elapsed() > MAX_PENDING_DURATION)
.map(|(id, _)| id.clone())
.collect()
};
for file_id in stale_files {
tracing::warn!("Forcing flush of stale pending file: {}", file_id);
let _ = self.flush_tx.try_send(FlushCommand::FlushFile(file_id));
}
}
}
}
// ─── Port implementation ─────────────────────────────────────────────────────
#[async_trait]
impl WriteBehindCachePort for WriteBehindCache {
fn is_eligible_size(&self, size: usize) -> bool {
WriteBehindCache::is_eligible(size)
}
async fn put_pending(
&self,
file_id: String,
content: Bytes,
target_path: PathBuf,
) -> Result<bool, DomainError> {
self.put_pending(file_id, content, target_path)
.await
.map_err(DomainError::from)
}
async fn get_pending(&self, file_id: &str) -> Option<Bytes> {
self.get_pending(file_id).await
}
async fn is_pending(&self, file_id: &str) -> bool {
self.is_pending(file_id).await
}
async fn force_flush(&self, file_id: &str) -> Result<(), DomainError> {
self.force_flush(file_id).await.map_err(DomainError::from)
}
async fn flush_all(&self) -> Result<(), DomainError> {
self.flush_all().await.map_err(DomainError::from)
}
async fn shutdown(&self) -> Result<(), DomainError> {
self.shutdown().await.map_err(DomainError::from)
}
async fn get_stats(&self) -> WriteBehindStatsDto {
let stats = self.get_stats().await;
WriteBehindStatsDto {
pending_count: stats.pending_count,
pending_bytes: stats.pending_bytes,
total_writes: stats.total_writes,
total_bytes_written: stats.total_bytes_written,
cache_hits: stats.cache_hits,
avg_flush_time_us: stats.avg_flush_time_us,
}
}
}
impl Default for WriteBehindCache {
fn default() -> Self {
// Note: This creates a non-Arc version, prefer using new()
let (flush_tx, _) = mpsc::channel(1);
Self {
pending: Arc::new(RwLock::new(HashMap::new())),
current_size: Arc::new(RwLock::new(0)),
flush_tx,
stats: Arc::new(RwLock::new(WriteBehindStats::default())),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::TempDir;
#[tokio::test]
async fn test_write_behind_basic() {
let cache = WriteBehindCache::new();
let temp_dir = TempDir::new().unwrap();
let target = temp_dir.path().join("test.txt");
let content = Bytes::from("Hello, World!");
// Put in cache
let cached = cache
.put_pending("test-id".to_string(), content.clone(), target.clone())
.await
.unwrap();
assert!(cached);
assert!(cache.is_pending("test-id").await);
// Should be readable from cache
let cached_content = cache.get_pending("test-id").await.unwrap();
assert_eq!(cached_content, content);
// Force flush
cache.force_flush("test-id").await.unwrap();
// Should no longer be pending
assert!(!cache.is_pending("test-id").await);
// File should exist on disk
assert!(target.exists());
let disk_content = std::fs::read(&target).unwrap();
assert_eq!(disk_content, content.as_ref());
}
#[tokio::test]
async fn test_eligibility() {
// 500KB should be eligible
assert!(WriteBehindCache::is_eligible(500 * 1024));
// 1MB exactly should be eligible
assert!(WriteBehindCache::is_eligible(1024 * 1024));
// Over 1MB should not be eligible
assert!(!WriteBehindCache::is_eligible(1024 * 1024 + 1));
}
}
+71 -53
View File
@@ -1,33 +1,33 @@
use std::io::{Cursor, Read, Write};
use zip::{ZipWriter, write::SimpleFileOptions};
use thiserror::Error;
use tracing::*;
use async_trait::async_trait;
use crate::{
application::dtos::file_dto::FileDto,
application::dtos::folder_dto::FolderDto,
application::ports::inbound::FolderUseCase,
application::ports::file_ports::FileRetrievalUseCase,
application::ports::inbound::FolderUseCase,
application::ports::zip_ports::ZipPort,
common::errors::{Result, DomainError, ErrorKind},
common::errors::{DomainError, ErrorKind, Result},
};
use async_trait::async_trait;
use std::io::{Cursor, Read, Write};
use std::sync::Arc;
use thiserror::Error;
use tracing::*;
use zip::{ZipWriter, write::SimpleFileOptions};
/// Error related to ZIP file creation
#[derive(Debug, Error)]
pub enum ZipError {
#[error("IO error: {0}")]
IoError(#[from] std::io::Error),
#[error("ZIP error: {0}")]
ZipError(#[from] zip::result::ZipError),
#[error("Error reading file: {0}")]
FileReadError(String),
#[error("Error getting folder contents: {0}")]
FolderContentsError(String),
#[error("Folder not found: {0}")]
FolderNotFound(String),
}
@@ -54,18 +54,24 @@ pub struct ZipService {
impl ZipService {
/// Creates a new instance of the ZIP service with a reference to the file service
pub fn new(file_service: Arc<dyn FileRetrievalUseCase>, folder_service: Arc<dyn FolderUseCase>) -> Self {
pub fn new(
file_service: Arc<dyn FileRetrievalUseCase>,
folder_service: Arc<dyn FolderUseCase>,
) -> Self {
Self {
file_service,
folder_service,
}
}
/// Creates a ZIP file with the contents of a folder and all its subfolders
/// Returns the ZIP bytes
pub async fn create_folder_zip(&self, folder_id: &str, folder_name: &str) -> Result<Vec<u8>> {
info!("Creating ZIP for folder: {} (ID: {})", folder_name, folder_id);
info!(
"Creating ZIP for folder: {} (ID: {})",
folder_name, folder_id
);
// Verify if the folder exists
let folder = match self.folder_service.get_folder(folder_id).await {
Ok(folder) => folder,
@@ -74,31 +80,32 @@ impl ZipService {
return Err(ZipError::FolderNotFound(folder_id.to_string()).into());
}
};
// Create an in-memory buffer for the ZIP
let buf = Cursor::new(Vec::new());
let mut zip = ZipWriter::new(buf);
// Set compression options
let options = SimpleFileOptions::default()
.compression_method(zip::CompressionMethod::Deflated)
.unix_permissions(0o755);
// Object to track processed folders and avoid cycles
let mut processed_folders = std::collections::HashSet::new();
// Process the root folder and build the ZIP
self.process_folder_recursively(
&mut zip,
&folder,
folder_name,
&options,
&mut processed_folders
).await?;
&mut processed_folders,
)
.await?;
// Finalize the ZIP and get the bytes
let mut zip_buf = zip.finish()?;
let mut bytes = Vec::new();
match zip_buf.read_to_end(&mut bytes) {
Ok(_) => Ok(bytes),
@@ -108,7 +115,7 @@ impl ZipService {
}
}
}
// Alternative implementation to avoid recursion in async
async fn process_folder_recursively(
&self,
@@ -116,31 +123,31 @@ impl ZipService {
folder: &FolderDto,
path: &str,
options: &SimpleFileOptions,
processed_folders: &mut std::collections::HashSet<String>
processed_folders: &mut std::collections::HashSet<String>,
) -> Result<()> {
// Structure to represent pending work
struct PendingFolder {
folder: FolderDto,
path: String,
}
// Work queue for iterative processing
let mut work_queue = vec![PendingFolder {
folder: folder.clone(),
path: path.to_string(),
}];
// Process the queue while there are elements
while let Some(current) = work_queue.pop() {
let folder_id = current.folder.id.to_string();
// Avoid cycles
if processed_folders.contains(&folder_id) {
continue;
}
processed_folders.insert(folder_id.clone());
// Create the directory entry in the ZIP
let folder_path = format!("{}/", current.path);
match zip.add_directory(&folder_path, *options) {
@@ -150,30 +157,39 @@ impl ZipService {
// Continue even if creating the directory fails (it could be a duplicate)
}
}
// Add files from the folder to the ZIP
let files = match self.file_service.list_files(Some(&folder_id)).await {
Ok(files) => files,
Err(e) => {
error!("Error listing files in folder {}: {}", folder_id, e);
return Err(ZipError::FolderContentsError(format!("Error listing files: {}", e)).into());
return Err(ZipError::FolderContentsError(format!(
"Error listing files: {}",
e
))
.into());
}
};
// Add each file to the ZIP
for file in files {
self.add_file_to_zip(zip, &file, &folder_path, options).await?;
self.add_file_to_zip(zip, &file, &folder_path, options)
.await?;
}
// Process subfolders
let subfolders = match self.folder_service.list_folders(Some(&folder_id)).await {
Ok(folders) => folders,
Err(e) => {
error!("Error listing subfolders in {}: {}", folder_id, e);
return Err(ZipError::FolderContentsError(format!("Error listing subfolders: {}", e)).into());
return Err(ZipError::FolderContentsError(format!(
"Error listing subfolders: {}",
e
))
.into());
}
};
// Add subfolders to the queue
for subfolder in subfolders {
let subfolder_path = format!("{}/{}", current.path, subfolder.name);
@@ -183,10 +199,10 @@ impl ZipService {
});
}
}
Ok(())
}
// Adds a file to the ZIP
async fn add_file_to_zip(
&self,
@@ -197,29 +213,31 @@ impl ZipService {
) -> Result<()> {
let file_path = format!("{}{}", folder_path, file.name);
info!("Adding file to ZIP: {}", file_path);
// Get the file content
let file_id = file.id.to_string();
let content = match self.file_service.get_file_content(&file_id).await {
Ok(content) => content,
Err(e) => {
error!("Error reading file content {}: {}", file_id, e);
return Err(ZipError::FileReadError(format!("Error reading file {}: {}", file_id, e)).into());
return Err(ZipError::FileReadError(format!(
"Error reading file {}: {}",
file_id, e
))
.into());
}
};
// Write file to the ZIP
match zip.start_file_from_path(std::path::Path::new(&file_path), *options) {
Ok(_) => {
match zip.write_all(&content) {
Ok(_) => {
debug!("File added to ZIP: {}", file_path);
Ok(())
},
Err(e) => {
error!("Error writing file content {}: {}", file_path, e);
Err(ZipError::IoError(e).into())
}
Ok(_) => match zip.write_all(&content) {
Ok(_) => {
debug!("File added to ZIP: {}", file_path);
Ok(())
}
Err(e) => {
error!("Error writing file content {}: {}", file_path, e);
Err(ZipError::IoError(e).into())
}
},
Err(e) => {
@@ -241,4 +259,4 @@ impl ZipPort for ZipService {
) -> std::result::Result<Vec<u8>, DomainError> {
self.create_folder_zip(folder_id, folder_name).await
}
}
}