style: apply cargo fmt to entire codebase
Standardize code formatting across all 173 Rust source files using rustfmt. No functional changes - purely cosmetic. This establishes a consistent code style baseline for the project going forward.
This commit is contained in:
@@ -1,298 +1,467 @@
|
||||
//! Calendar Storage Adapter
|
||||
//!
|
||||
//! This adapter implements the `CalendarStoragePort` application port using
|
||||
//! the `CalendarRepository` and `CalendarEventRepository` domain repositories.
|
||||
//! It bridges the gap between the application layer and the infrastructure layer.
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::collections::HashMap;
|
||||
use async_trait::async_trait;
|
||||
use chrono::{DateTime, Utc};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::application::dtos::calendar_dto::{
|
||||
CalendarDto, CalendarEventDto, CreateCalendarDto, UpdateCalendarDto,
|
||||
CreateEventDto, UpdateEventDto, CreateEventICalDto
|
||||
};
|
||||
use crate::application::ports::calendar_ports::CalendarStoragePort;
|
||||
use crate::common::errors::{DomainError, ErrorKind};
|
||||
use crate::domain::entities::calendar::Calendar;
|
||||
use crate::domain::entities::calendar_event::CalendarEvent;
|
||||
use crate::domain::repositories::calendar_repository::CalendarRepository;
|
||||
use crate::domain::repositories::calendar_event_repository::CalendarEventRepository;
|
||||
|
||||
/// Adapter that implements CalendarStoragePort using domain repositories
|
||||
pub struct CalendarStorageAdapter {
|
||||
calendar_repository: Arc<dyn CalendarRepository>,
|
||||
event_repository: Arc<dyn CalendarEventRepository>,
|
||||
}
|
||||
|
||||
impl CalendarStorageAdapter {
|
||||
/// Creates a new CalendarStorageAdapter with the given repositories
|
||||
pub fn new(
|
||||
calendar_repository: Arc<dyn CalendarRepository>,
|
||||
event_repository: Arc<dyn CalendarEventRepository>,
|
||||
) -> Self {
|
||||
Self {
|
||||
calendar_repository,
|
||||
event_repository,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl CalendarStoragePort for CalendarStorageAdapter {
|
||||
// Calendar operations
|
||||
|
||||
async fn create_calendar(&self, dto: CreateCalendarDto, owner_id: &str) -> Result<CalendarDto, DomainError> {
|
||||
let calendar = Calendar::new(
|
||||
dto.name,
|
||||
owner_id.to_string(),
|
||||
dto.description,
|
||||
dto.color,
|
||||
)?;
|
||||
|
||||
let created = self.calendar_repository.create_calendar(calendar).await?;
|
||||
Ok(CalendarDto::from(created))
|
||||
}
|
||||
|
||||
async fn update_calendar(&self, calendar_id: &str, update: UpdateCalendarDto) -> Result<CalendarDto, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
let mut calendar = self.calendar_repository.find_calendar_by_id(&uuid).await?;
|
||||
|
||||
if let Some(name) = update.name {
|
||||
calendar.update_name(name)?;
|
||||
}
|
||||
if let Some(description) = update.description {
|
||||
calendar.update_description(Some(description));
|
||||
}
|
||||
if let Some(color) = update.color {
|
||||
calendar.update_color(Some(color))?;
|
||||
}
|
||||
|
||||
let updated = self.calendar_repository.update_calendar(calendar).await?;
|
||||
Ok(CalendarDto::from(updated))
|
||||
}
|
||||
|
||||
async fn delete_calendar(&self, calendar_id: &str) -> Result<(), DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
// First delete all events in the calendar
|
||||
self.event_repository.delete_all_events_in_calendar(&uuid).await?;
|
||||
|
||||
// Then delete the calendar itself
|
||||
self.calendar_repository.delete_calendar(&uuid).await
|
||||
}
|
||||
|
||||
async fn get_calendar(&self, calendar_id: &str) -> Result<CalendarDto, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
let calendar = self.calendar_repository.find_calendar_by_id(&uuid).await?;
|
||||
Ok(CalendarDto::from(calendar))
|
||||
}
|
||||
|
||||
async fn list_calendars_by_owner(&self, owner_id: &str) -> Result<Vec<CalendarDto>, DomainError> {
|
||||
let calendars = self.calendar_repository.list_calendars_by_owner(owner_id).await?;
|
||||
Ok(calendars.into_iter().map(CalendarDto::from).collect())
|
||||
}
|
||||
|
||||
async fn list_calendars_shared_with_user(&self, user_id: &str) -> Result<Vec<CalendarDto>, DomainError> {
|
||||
let calendars = self.calendar_repository.list_calendars_shared_with_user(user_id).await?;
|
||||
Ok(calendars.into_iter().map(CalendarDto::from).collect())
|
||||
}
|
||||
|
||||
async fn list_public_calendars(&self, limit: i64, offset: i64) -> Result<Vec<CalendarDto>, DomainError> {
|
||||
let calendars = self.calendar_repository.list_public_calendars(limit, offset).await?;
|
||||
Ok(calendars.into_iter().map(CalendarDto::from).collect())
|
||||
}
|
||||
|
||||
async fn check_calendar_access(&self, calendar_id: &str, user_id: &str) -> Result<bool, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
self.calendar_repository.user_has_calendar_access(&uuid, user_id).await
|
||||
}
|
||||
|
||||
// Calendar sharing
|
||||
|
||||
async fn share_calendar(&self, calendar_id: &str, user_id: &str, access_level: &str) -> Result<(), DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
self.calendar_repository.share_calendar(&uuid, user_id, access_level).await
|
||||
}
|
||||
|
||||
async fn remove_calendar_sharing(&self, calendar_id: &str, user_id: &str) -> Result<(), DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
self.calendar_repository.remove_calendar_sharing(&uuid, user_id).await
|
||||
}
|
||||
|
||||
async fn get_calendar_shares(&self, calendar_id: &str) -> Result<Vec<(String, String)>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
self.calendar_repository.get_calendar_shares(&uuid).await
|
||||
}
|
||||
|
||||
// Calendar properties
|
||||
|
||||
async fn set_calendar_property(&self, calendar_id: &str, property_name: &str, property_value: &str) -> Result<(), DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
self.calendar_repository.set_calendar_property(&uuid, property_name, property_value).await
|
||||
}
|
||||
|
||||
async fn get_calendar_property(&self, calendar_id: &str, property_name: &str) -> Result<Option<String>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
self.calendar_repository.get_calendar_property(&uuid, property_name).await
|
||||
}
|
||||
|
||||
async fn get_calendar_properties(&self, calendar_id: &str) -> Result<HashMap<String, String>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
self.calendar_repository.get_calendar_properties(&uuid).await
|
||||
}
|
||||
|
||||
// Event operations
|
||||
|
||||
async fn create_event(&self, dto: CreateEventDto) -> Result<CalendarEventDto, DomainError> {
|
||||
let calendar_id = Uuid::parse_str(&dto.calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Event", "Invalid calendar ID format"))?;
|
||||
|
||||
// Verify calendar exists and user has access
|
||||
let _calendar = self.calendar_repository.find_calendar_by_id(&calendar_id).await?;
|
||||
|
||||
// Generate basic iCal data
|
||||
let ical_data = format!(
|
||||
"BEGIN:VCALENDAR\nVERSION:2.0\nPRODID:-//OxiCloud//EN\nBEGIN:VEVENT\nUID:{}@oxicloud\nDTSTAMP:{}\nDTSTART:{}\nDTEND:{}\nSUMMARY:{}\nEND:VEVENT\nEND:VCALENDAR",
|
||||
uuid::Uuid::new_v4(),
|
||||
chrono::Utc::now().format("%Y%m%dT%H%M%SZ"),
|
||||
dto.start_time.format("%Y%m%dT%H%M%SZ"),
|
||||
dto.end_time.format("%Y%m%dT%H%M%SZ"),
|
||||
dto.summary
|
||||
);
|
||||
|
||||
let event = CalendarEvent::new(
|
||||
calendar_id,
|
||||
dto.summary,
|
||||
dto.description,
|
||||
dto.location,
|
||||
dto.start_time,
|
||||
dto.end_time,
|
||||
dto.all_day.unwrap_or(false),
|
||||
dto.rrule,
|
||||
ical_data,
|
||||
)?;
|
||||
|
||||
let created = self.event_repository.create_event(event).await?;
|
||||
Ok(CalendarEventDto::from(created))
|
||||
}
|
||||
|
||||
async fn create_event_from_ical(&self, dto: CreateEventICalDto) -> Result<CalendarEventDto, DomainError> {
|
||||
let calendar_id = Uuid::parse_str(&dto.calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Event", "Invalid calendar ID format"))?;
|
||||
|
||||
// Verify calendar exists
|
||||
let _calendar = self.calendar_repository.find_calendar_by_id(&calendar_id).await?;
|
||||
|
||||
// Parse iCal data and create event
|
||||
let event = CalendarEvent::from_ical(calendar_id, dto.ical_data.clone())?;
|
||||
|
||||
let created = self.event_repository.create_event(event).await?;
|
||||
Ok(CalendarEventDto::from(created))
|
||||
}
|
||||
|
||||
async fn update_event(&self, event_id: &str, update: UpdateEventDto) -> Result<CalendarEventDto, DomainError> {
|
||||
let uuid = Uuid::parse_str(event_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Event", "Invalid event ID format"))?;
|
||||
|
||||
let mut event = self.event_repository.find_event_by_id(&uuid).await?;
|
||||
|
||||
if let Some(summary) = update.summary {
|
||||
event.update_summary(summary)?;
|
||||
}
|
||||
if let Some(description) = update.description {
|
||||
event.update_description(Some(description));
|
||||
}
|
||||
if let Some(location) = update.location {
|
||||
event.update_location(Some(location));
|
||||
}
|
||||
if let Some(start_time) = update.start_time {
|
||||
if let Some(end_time) = update.end_time {
|
||||
event.update_time_range(start_time, end_time)?;
|
||||
} else {
|
||||
event.update_time_range(start_time, *event.end_time())?;
|
||||
}
|
||||
} else if let Some(end_time) = update.end_time {
|
||||
event.update_time_range(*event.start_time(), end_time)?;
|
||||
}
|
||||
if let Some(all_day) = update.all_day {
|
||||
event.update_all_day(all_day);
|
||||
}
|
||||
if let Some(rrule) = update.rrule {
|
||||
event.update_rrule(Some(rrule))?;
|
||||
}
|
||||
|
||||
let updated = self.event_repository.update_event(event).await?;
|
||||
Ok(CalendarEventDto::from(updated))
|
||||
}
|
||||
|
||||
async fn delete_event(&self, event_id: &str) -> Result<(), DomainError> {
|
||||
let uuid = Uuid::parse_str(event_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Event", "Invalid event ID format"))?;
|
||||
|
||||
self.event_repository.delete_event(&uuid).await
|
||||
}
|
||||
|
||||
async fn get_event(&self, event_id: &str) -> Result<CalendarEventDto, DomainError> {
|
||||
let uuid = Uuid::parse_str(event_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Event", "Invalid event ID format"))?;
|
||||
|
||||
let event = self.event_repository.find_event_by_id(&uuid).await?;
|
||||
Ok(CalendarEventDto::from(event))
|
||||
}
|
||||
|
||||
async fn list_events_by_calendar(&self, calendar_id: &str) -> Result<Vec<CalendarEventDto>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
let events = self.event_repository.list_events_by_calendar(&uuid).await?;
|
||||
Ok(events.into_iter().map(CalendarEventDto::from).collect())
|
||||
}
|
||||
|
||||
async fn list_events_by_calendar_paginated(&self, calendar_id: &str, limit: i64, offset: i64) -> Result<Vec<CalendarEventDto>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
let events = self.event_repository.list_events_by_calendar_paginated(&uuid, limit, offset).await?;
|
||||
Ok(events.into_iter().map(CalendarEventDto::from).collect())
|
||||
}
|
||||
|
||||
async fn get_events_in_time_range(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
start: &DateTime<Utc>,
|
||||
end: &DateTime<Utc>
|
||||
) -> Result<Vec<CalendarEventDto>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id)
|
||||
.map_err(|_| DomainError::new(ErrorKind::InvalidInput, "Calendar", "Invalid calendar ID format"))?;
|
||||
|
||||
let events = self.event_repository.get_events_in_time_range(&uuid, start, end).await?;
|
||||
Ok(events.into_iter().map(CalendarEventDto::from).collect())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
// Tests would go here using mock repositories
|
||||
}
|
||||
//! Calendar Storage Adapter
|
||||
//!
|
||||
//! This adapter implements the `CalendarStoragePort` application port using
|
||||
//! the `CalendarRepository` and `CalendarEventRepository` domain repositories.
|
||||
//! It bridges the gap between the application layer and the infrastructure layer.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use chrono::{DateTime, Utc};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::application::dtos::calendar_dto::{
|
||||
CalendarDto, CalendarEventDto, CreateCalendarDto, CreateEventDto, CreateEventICalDto,
|
||||
UpdateCalendarDto, UpdateEventDto,
|
||||
};
|
||||
use crate::application::ports::calendar_ports::CalendarStoragePort;
|
||||
use crate::common::errors::{DomainError, ErrorKind};
|
||||
use crate::domain::entities::calendar::Calendar;
|
||||
use crate::domain::entities::calendar_event::CalendarEvent;
|
||||
use crate::domain::repositories::calendar_event_repository::CalendarEventRepository;
|
||||
use crate::domain::repositories::calendar_repository::CalendarRepository;
|
||||
|
||||
/// Adapter that implements CalendarStoragePort using domain repositories
|
||||
pub struct CalendarStorageAdapter {
|
||||
calendar_repository: Arc<dyn CalendarRepository>,
|
||||
event_repository: Arc<dyn CalendarEventRepository>,
|
||||
}
|
||||
|
||||
impl CalendarStorageAdapter {
|
||||
/// Creates a new CalendarStorageAdapter with the given repositories
|
||||
pub fn new(
|
||||
calendar_repository: Arc<dyn CalendarRepository>,
|
||||
event_repository: Arc<dyn CalendarEventRepository>,
|
||||
) -> Self {
|
||||
Self {
|
||||
calendar_repository,
|
||||
event_repository,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl CalendarStoragePort for CalendarStorageAdapter {
|
||||
// Calendar operations
|
||||
|
||||
async fn create_calendar(
|
||||
&self,
|
||||
dto: CreateCalendarDto,
|
||||
owner_id: &str,
|
||||
) -> Result<CalendarDto, DomainError> {
|
||||
let calendar = Calendar::new(dto.name, owner_id.to_string(), dto.description, dto.color)?;
|
||||
|
||||
let created = self.calendar_repository.create_calendar(calendar).await?;
|
||||
Ok(CalendarDto::from(created))
|
||||
}
|
||||
|
||||
async fn update_calendar(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
update: UpdateCalendarDto,
|
||||
) -> Result<CalendarDto, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
let mut calendar = self.calendar_repository.find_calendar_by_id(&uuid).await?;
|
||||
|
||||
if let Some(name) = update.name {
|
||||
calendar.update_name(name)?;
|
||||
}
|
||||
if let Some(description) = update.description {
|
||||
calendar.update_description(Some(description));
|
||||
}
|
||||
if let Some(color) = update.color {
|
||||
calendar.update_color(Some(color))?;
|
||||
}
|
||||
|
||||
let updated = self.calendar_repository.update_calendar(calendar).await?;
|
||||
Ok(CalendarDto::from(updated))
|
||||
}
|
||||
|
||||
async fn delete_calendar(&self, calendar_id: &str) -> Result<(), DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
// First delete all events in the calendar
|
||||
self.event_repository
|
||||
.delete_all_events_in_calendar(&uuid)
|
||||
.await?;
|
||||
|
||||
// Then delete the calendar itself
|
||||
self.calendar_repository.delete_calendar(&uuid).await
|
||||
}
|
||||
|
||||
async fn get_calendar(&self, calendar_id: &str) -> Result<CalendarDto, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
let calendar = self.calendar_repository.find_calendar_by_id(&uuid).await?;
|
||||
Ok(CalendarDto::from(calendar))
|
||||
}
|
||||
|
||||
async fn list_calendars_by_owner(
|
||||
&self,
|
||||
owner_id: &str,
|
||||
) -> Result<Vec<CalendarDto>, DomainError> {
|
||||
let calendars = self
|
||||
.calendar_repository
|
||||
.list_calendars_by_owner(owner_id)
|
||||
.await?;
|
||||
Ok(calendars.into_iter().map(CalendarDto::from).collect())
|
||||
}
|
||||
|
||||
async fn list_calendars_shared_with_user(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<Vec<CalendarDto>, DomainError> {
|
||||
let calendars = self
|
||||
.calendar_repository
|
||||
.list_calendars_shared_with_user(user_id)
|
||||
.await?;
|
||||
Ok(calendars.into_iter().map(CalendarDto::from).collect())
|
||||
}
|
||||
|
||||
async fn list_public_calendars(
|
||||
&self,
|
||||
limit: i64,
|
||||
offset: i64,
|
||||
) -> Result<Vec<CalendarDto>, DomainError> {
|
||||
let calendars = self
|
||||
.calendar_repository
|
||||
.list_public_calendars(limit, offset)
|
||||
.await?;
|
||||
Ok(calendars.into_iter().map(CalendarDto::from).collect())
|
||||
}
|
||||
|
||||
async fn check_calendar_access(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
self.calendar_repository
|
||||
.user_has_calendar_access(&uuid, user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
// Calendar sharing
|
||||
|
||||
async fn share_calendar(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
user_id: &str,
|
||||
access_level: &str,
|
||||
) -> Result<(), DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
self.calendar_repository
|
||||
.share_calendar(&uuid, user_id, access_level)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn remove_calendar_sharing(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<(), DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
self.calendar_repository
|
||||
.remove_calendar_sharing(&uuid, user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_calendar_shares(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
) -> Result<Vec<(String, String)>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
self.calendar_repository.get_calendar_shares(&uuid).await
|
||||
}
|
||||
|
||||
// Calendar properties
|
||||
|
||||
async fn set_calendar_property(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
property_name: &str,
|
||||
property_value: &str,
|
||||
) -> Result<(), DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
self.calendar_repository
|
||||
.set_calendar_property(&uuid, property_name, property_value)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_calendar_property(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
property_name: &str,
|
||||
) -> Result<Option<String>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
self.calendar_repository
|
||||
.get_calendar_property(&uuid, property_name)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_calendar_properties(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
) -> Result<HashMap<String, String>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
self.calendar_repository
|
||||
.get_calendar_properties(&uuid)
|
||||
.await
|
||||
}
|
||||
|
||||
// Event operations
|
||||
|
||||
async fn create_event(&self, dto: CreateEventDto) -> Result<CalendarEventDto, DomainError> {
|
||||
let calendar_id = Uuid::parse_str(&dto.calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Event",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
// Verify calendar exists and user has access
|
||||
let _calendar = self
|
||||
.calendar_repository
|
||||
.find_calendar_by_id(&calendar_id)
|
||||
.await?;
|
||||
|
||||
// Generate basic iCal data
|
||||
let ical_data = format!(
|
||||
"BEGIN:VCALENDAR\nVERSION:2.0\nPRODID:-//OxiCloud//EN\nBEGIN:VEVENT\nUID:{}@oxicloud\nDTSTAMP:{}\nDTSTART:{}\nDTEND:{}\nSUMMARY:{}\nEND:VEVENT\nEND:VCALENDAR",
|
||||
uuid::Uuid::new_v4(),
|
||||
chrono::Utc::now().format("%Y%m%dT%H%M%SZ"),
|
||||
dto.start_time.format("%Y%m%dT%H%M%SZ"),
|
||||
dto.end_time.format("%Y%m%dT%H%M%SZ"),
|
||||
dto.summary
|
||||
);
|
||||
|
||||
let event = CalendarEvent::new(
|
||||
calendar_id,
|
||||
dto.summary,
|
||||
dto.description,
|
||||
dto.location,
|
||||
dto.start_time,
|
||||
dto.end_time,
|
||||
dto.all_day.unwrap_or(false),
|
||||
dto.rrule,
|
||||
ical_data,
|
||||
)?;
|
||||
|
||||
let created = self.event_repository.create_event(event).await?;
|
||||
Ok(CalendarEventDto::from(created))
|
||||
}
|
||||
|
||||
async fn create_event_from_ical(
|
||||
&self,
|
||||
dto: CreateEventICalDto,
|
||||
) -> Result<CalendarEventDto, DomainError> {
|
||||
let calendar_id = Uuid::parse_str(&dto.calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Event",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
// Verify calendar exists
|
||||
let _calendar = self
|
||||
.calendar_repository
|
||||
.find_calendar_by_id(&calendar_id)
|
||||
.await?;
|
||||
|
||||
// Parse iCal data and create event
|
||||
let event = CalendarEvent::from_ical(calendar_id, dto.ical_data.clone())?;
|
||||
|
||||
let created = self.event_repository.create_event(event).await?;
|
||||
Ok(CalendarEventDto::from(created))
|
||||
}
|
||||
|
||||
async fn update_event(
|
||||
&self,
|
||||
event_id: &str,
|
||||
update: UpdateEventDto,
|
||||
) -> Result<CalendarEventDto, DomainError> {
|
||||
let uuid = Uuid::parse_str(event_id).map_err(|_| {
|
||||
DomainError::new(ErrorKind::InvalidInput, "Event", "Invalid event ID format")
|
||||
})?;
|
||||
|
||||
let mut event = self.event_repository.find_event_by_id(&uuid).await?;
|
||||
|
||||
if let Some(summary) = update.summary {
|
||||
event.update_summary(summary)?;
|
||||
}
|
||||
if let Some(description) = update.description {
|
||||
event.update_description(Some(description));
|
||||
}
|
||||
if let Some(location) = update.location {
|
||||
event.update_location(Some(location));
|
||||
}
|
||||
if let Some(start_time) = update.start_time {
|
||||
if let Some(end_time) = update.end_time {
|
||||
event.update_time_range(start_time, end_time)?;
|
||||
} else {
|
||||
event.update_time_range(start_time, *event.end_time())?;
|
||||
}
|
||||
} else if let Some(end_time) = update.end_time {
|
||||
event.update_time_range(*event.start_time(), end_time)?;
|
||||
}
|
||||
if let Some(all_day) = update.all_day {
|
||||
event.update_all_day(all_day);
|
||||
}
|
||||
if let Some(rrule) = update.rrule {
|
||||
event.update_rrule(Some(rrule))?;
|
||||
}
|
||||
|
||||
let updated = self.event_repository.update_event(event).await?;
|
||||
Ok(CalendarEventDto::from(updated))
|
||||
}
|
||||
|
||||
async fn delete_event(&self, event_id: &str) -> Result<(), DomainError> {
|
||||
let uuid = Uuid::parse_str(event_id).map_err(|_| {
|
||||
DomainError::new(ErrorKind::InvalidInput, "Event", "Invalid event ID format")
|
||||
})?;
|
||||
|
||||
self.event_repository.delete_event(&uuid).await
|
||||
}
|
||||
|
||||
async fn get_event(&self, event_id: &str) -> Result<CalendarEventDto, DomainError> {
|
||||
let uuid = Uuid::parse_str(event_id).map_err(|_| {
|
||||
DomainError::new(ErrorKind::InvalidInput, "Event", "Invalid event ID format")
|
||||
})?;
|
||||
|
||||
let event = self.event_repository.find_event_by_id(&uuid).await?;
|
||||
Ok(CalendarEventDto::from(event))
|
||||
}
|
||||
|
||||
async fn list_events_by_calendar(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
) -> Result<Vec<CalendarEventDto>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
let events = self.event_repository.list_events_by_calendar(&uuid).await?;
|
||||
Ok(events.into_iter().map(CalendarEventDto::from).collect())
|
||||
}
|
||||
|
||||
async fn list_events_by_calendar_paginated(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
limit: i64,
|
||||
offset: i64,
|
||||
) -> Result<Vec<CalendarEventDto>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
let events = self
|
||||
.event_repository
|
||||
.list_events_by_calendar_paginated(&uuid, limit, offset)
|
||||
.await?;
|
||||
Ok(events.into_iter().map(CalendarEventDto::from).collect())
|
||||
}
|
||||
|
||||
async fn get_events_in_time_range(
|
||||
&self,
|
||||
calendar_id: &str,
|
||||
start: &DateTime<Utc>,
|
||||
end: &DateTime<Utc>,
|
||||
) -> Result<Vec<CalendarEventDto>, DomainError> {
|
||||
let uuid = Uuid::parse_str(calendar_id).map_err(|_| {
|
||||
DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Calendar",
|
||||
"Invalid calendar ID format",
|
||||
)
|
||||
})?;
|
||||
|
||||
let events = self
|
||||
.event_repository
|
||||
.get_events_in_time_range(&uuid, start, end)
|
||||
.await?;
|
||||
Ok(events.into_iter().map(CalendarEventDto::from).collect())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
// Tests would go here using mock repositories
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,125 +1,128 @@
|
||||
//! Infrastructure Error Adapters
|
||||
//!
|
||||
//! This module contains error conversion adapters for infrastructure-specific errors.
|
||||
//! These adapters bridge the gap between infrastructure errors (sqlx, serde_json, etc.)
|
||||
//! and domain errors, keeping the domain layer clean of infrastructure knowledge.
|
||||
//!
|
||||
//! Following Clean Architecture principles, these conversions are placed in the
|
||||
//! infrastructure layer rather than the common/domain layers.
|
||||
|
||||
use crate::domain::errors::{DomainError, ErrorKind};
|
||||
|
||||
/// Macro to create From implementations for infrastructure errors to DomainError.
|
||||
///
|
||||
/// This macro is intended for use ONLY within the infrastructure layer.
|
||||
/// The domain layer should not depend on specific infrastructure error types.
|
||||
///
|
||||
/// # Example
|
||||
///
|
||||
/// ```ignore
|
||||
/// // In infrastructure code:
|
||||
/// impl_infra_error_to_domain!(serde_json::Error, "Serialization");
|
||||
/// impl_infra_error_to_domain!(sqlx::Error, "Database");
|
||||
/// ```
|
||||
#[macro_export]
|
||||
macro_rules! impl_infra_error_to_domain {
|
||||
($error_type:ty, $entity_type:expr) => {
|
||||
impl From<$error_type> for $crate::domain::errors::DomainError {
|
||||
fn from(err: $error_type) -> Self {
|
||||
$crate::domain::errors::DomainError {
|
||||
kind: $crate::domain::errors::ErrorKind::InternalError,
|
||||
entity_type: $entity_type,
|
||||
entity_id: None,
|
||||
message: format!("{}", err),
|
||||
source: Some(Box::new(err)),
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
// Note: We intentionally DO NOT create global From implementations for sqlx::Error
|
||||
// or serde_json::Error here. Each repository/service should handle its own error
|
||||
// conversions with proper context. This prevents the domain from depending on
|
||||
// infrastructure error types.
|
||||
|
||||
/// Helper trait for converting infrastructure errors to DomainError with context.
|
||||
///
|
||||
/// This trait provides a more explicit way to convert infrastructure errors
|
||||
/// to domain errors, requiring the caller to provide context about the entity
|
||||
/// being operated on.
|
||||
pub trait IntoDomainError {
|
||||
/// Convert the error to a DomainError with the given entity type context.
|
||||
fn into_domain_error(self, entity_type: &'static str) -> DomainError;
|
||||
}
|
||||
|
||||
impl IntoDomainError for std::io::Error {
|
||||
fn into_domain_error(self, entity_type: &'static str) -> DomainError {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
entity_type,
|
||||
format!("IO error: {}", self),
|
||||
).with_source(self)
|
||||
}
|
||||
}
|
||||
|
||||
impl IntoDomainError for serde_json::Error {
|
||||
fn into_domain_error(self, entity_type: &'static str) -> DomainError {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
entity_type,
|
||||
format!("Serialization error: {}", self),
|
||||
).with_source(self)
|
||||
}
|
||||
}
|
||||
|
||||
impl IntoDomainError for sqlx::Error {
|
||||
fn into_domain_error(self, entity_type: &'static str) -> DomainError {
|
||||
match &self {
|
||||
sqlx::Error::RowNotFound => {
|
||||
DomainError::not_found(entity_type, "Record not found")
|
||||
}
|
||||
sqlx::Error::Database(db_err) => {
|
||||
// Handle specific PostgreSQL error codes
|
||||
if db_err.code().is_some_and(|c| c == "23505") {
|
||||
DomainError::already_exists(entity_type, "Record already exists")
|
||||
} else {
|
||||
DomainError::new(
|
||||
ErrorKind::DatabaseError,
|
||||
entity_type,
|
||||
format!("Database error: {}", db_err),
|
||||
).with_source(self)
|
||||
}
|
||||
}
|
||||
_ => DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
entity_type,
|
||||
format!("Database error: {}", self),
|
||||
).with_source(self)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_io_error_conversion() {
|
||||
let io_error = std::io::Error::new(std::io::ErrorKind::NotFound, "file not found");
|
||||
let domain_error = io_error.into_domain_error("File");
|
||||
|
||||
assert_eq!(domain_error.entity_type, "File");
|
||||
assert!(domain_error.message.contains("IO error"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_serde_json_error_conversion() {
|
||||
let json_str = "{ invalid json }";
|
||||
let serde_error: serde_json::Error = serde_json::from_str::<serde_json::Value>(json_str).unwrap_err();
|
||||
let domain_error = serde_error.into_domain_error("Config");
|
||||
|
||||
assert_eq!(domain_error.entity_type, "Config");
|
||||
assert!(domain_error.message.contains("Serialization error"));
|
||||
}
|
||||
}
|
||||
//! Infrastructure Error Adapters
|
||||
//!
|
||||
//! This module contains error conversion adapters for infrastructure-specific errors.
|
||||
//! These adapters bridge the gap between infrastructure errors (sqlx, serde_json, etc.)
|
||||
//! and domain errors, keeping the domain layer clean of infrastructure knowledge.
|
||||
//!
|
||||
//! Following Clean Architecture principles, these conversions are placed in the
|
||||
//! infrastructure layer rather than the common/domain layers.
|
||||
|
||||
use crate::domain::errors::{DomainError, ErrorKind};
|
||||
|
||||
/// Macro to create From implementations for infrastructure errors to DomainError.
|
||||
///
|
||||
/// This macro is intended for use ONLY within the infrastructure layer.
|
||||
/// The domain layer should not depend on specific infrastructure error types.
|
||||
///
|
||||
/// # Example
|
||||
///
|
||||
/// ```ignore
|
||||
/// // In infrastructure code:
|
||||
/// impl_infra_error_to_domain!(serde_json::Error, "Serialization");
|
||||
/// impl_infra_error_to_domain!(sqlx::Error, "Database");
|
||||
/// ```
|
||||
#[macro_export]
|
||||
macro_rules! impl_infra_error_to_domain {
|
||||
($error_type:ty, $entity_type:expr) => {
|
||||
impl From<$error_type> for $crate::domain::errors::DomainError {
|
||||
fn from(err: $error_type) -> Self {
|
||||
$crate::domain::errors::DomainError {
|
||||
kind: $crate::domain::errors::ErrorKind::InternalError,
|
||||
entity_type: $entity_type,
|
||||
entity_id: None,
|
||||
message: format!("{}", err),
|
||||
source: Some(Box::new(err)),
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
// Note: We intentionally DO NOT create global From implementations for sqlx::Error
|
||||
// or serde_json::Error here. Each repository/service should handle its own error
|
||||
// conversions with proper context. This prevents the domain from depending on
|
||||
// infrastructure error types.
|
||||
|
||||
/// Helper trait for converting infrastructure errors to DomainError with context.
|
||||
///
|
||||
/// This trait provides a more explicit way to convert infrastructure errors
|
||||
/// to domain errors, requiring the caller to provide context about the entity
|
||||
/// being operated on.
|
||||
pub trait IntoDomainError {
|
||||
/// Convert the error to a DomainError with the given entity type context.
|
||||
fn into_domain_error(self, entity_type: &'static str) -> DomainError;
|
||||
}
|
||||
|
||||
impl IntoDomainError for std::io::Error {
|
||||
fn into_domain_error(self, entity_type: &'static str) -> DomainError {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
entity_type,
|
||||
format!("IO error: {}", self),
|
||||
)
|
||||
.with_source(self)
|
||||
}
|
||||
}
|
||||
|
||||
impl IntoDomainError for serde_json::Error {
|
||||
fn into_domain_error(self, entity_type: &'static str) -> DomainError {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
entity_type,
|
||||
format!("Serialization error: {}", self),
|
||||
)
|
||||
.with_source(self)
|
||||
}
|
||||
}
|
||||
|
||||
impl IntoDomainError for sqlx::Error {
|
||||
fn into_domain_error(self, entity_type: &'static str) -> DomainError {
|
||||
match &self {
|
||||
sqlx::Error::RowNotFound => DomainError::not_found(entity_type, "Record not found"),
|
||||
sqlx::Error::Database(db_err) => {
|
||||
// Handle specific PostgreSQL error codes
|
||||
if db_err.code().is_some_and(|c| c == "23505") {
|
||||
DomainError::already_exists(entity_type, "Record already exists")
|
||||
} else {
|
||||
DomainError::new(
|
||||
ErrorKind::DatabaseError,
|
||||
entity_type,
|
||||
format!("Database error: {}", db_err),
|
||||
)
|
||||
.with_source(self)
|
||||
}
|
||||
}
|
||||
_ => DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
entity_type,
|
||||
format!("Database error: {}", self),
|
||||
)
|
||||
.with_source(self),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_io_error_conversion() {
|
||||
let io_error = std::io::Error::new(std::io::ErrorKind::NotFound, "file not found");
|
||||
let domain_error = io_error.into_domain_error("File");
|
||||
|
||||
assert_eq!(domain_error.entity_type, "File");
|
||||
assert!(domain_error.message.contains("IO error"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_serde_json_error_conversion() {
|
||||
let json_str = "{ invalid json }";
|
||||
let serde_error: serde_json::Error =
|
||||
serde_json::from_str::<serde_json::Value>(json_str).unwrap_err();
|
||||
let domain_error = serde_error.into_domain_error("Config");
|
||||
|
||||
assert_eq!(domain_error.entity_type, "Config");
|
||||
assert!(domain_error.message.contains("Serialization error"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,16 +1,16 @@
|
||||
//! Infrastructure Adapters
|
||||
//!
|
||||
//! This module contains adapters that bridge the gap between domain repositories
|
||||
//! and application ports. These adapters implement the application layer ports
|
||||
//! using the infrastructure layer repositories.
|
||||
//!
|
||||
//! It also includes error adapters for converting infrastructure-specific errors
|
||||
//! to domain errors, following Clean Architecture principles.
|
||||
|
||||
pub mod calendar_storage_adapter;
|
||||
pub mod contact_storage_adapter;
|
||||
pub mod error_adapters;
|
||||
|
||||
pub use calendar_storage_adapter::CalendarStorageAdapter;
|
||||
pub use contact_storage_adapter::ContactStorageAdapter;
|
||||
pub use error_adapters::IntoDomainError;
|
||||
//! Infrastructure Adapters
|
||||
//!
|
||||
//! This module contains adapters that bridge the gap between domain repositories
|
||||
//! and application ports. These adapters implement the application layer ports
|
||||
//! using the infrastructure layer repositories.
|
||||
//!
|
||||
//! It also includes error adapters for converting infrastructure-specific errors
|
||||
//! to domain errors, following Clean Architecture principles.
|
||||
|
||||
pub mod calendar_storage_adapter;
|
||||
pub mod contact_storage_adapter;
|
||||
pub mod error_adapters;
|
||||
|
||||
pub use calendar_storage_adapter::CalendarStorageAdapter;
|
||||
pub use contact_storage_adapter::ContactStorageAdapter;
|
||||
pub use error_adapters::IntoDomainError;
|
||||
|
||||
@@ -1,21 +1,21 @@
|
||||
use std::sync::Arc;
|
||||
use anyhow::Result;
|
||||
use sqlx::PgPool;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::application::ports::auth_ports::TokenServicePort;
|
||||
use crate::application::services::auth_application_service::AuthApplicationService;
|
||||
use crate::application::services::folder_service::FolderService;
|
||||
use crate::infrastructure::repositories::{UserPgRepository, SessionPgRepository};
|
||||
use crate::infrastructure::services::password_hasher::Argon2PasswordHasher;
|
||||
use crate::infrastructure::services::jwt_service::JwtTokenService;
|
||||
use crate::infrastructure::services::oidc_service::OidcService;
|
||||
use crate::common::config::AppConfig;
|
||||
use crate::common::di::AuthServices;
|
||||
use crate::infrastructure::repositories::{SessionPgRepository, UserPgRepository};
|
||||
use crate::infrastructure::services::jwt_service::JwtTokenService;
|
||||
use crate::infrastructure::services::oidc_service::OidcService;
|
||||
use crate::infrastructure::services::password_hasher::Argon2PasswordHasher;
|
||||
|
||||
pub async fn create_auth_services(
|
||||
config: &AppConfig,
|
||||
config: &AppConfig,
|
||||
pool: Arc<PgPool>,
|
||||
folder_service: Option<Arc<FolderService>>
|
||||
folder_service: Option<Arc<FolderService>>,
|
||||
) -> Result<AuthServices> {
|
||||
// Create JWT token service (TokenServicePort implementation)
|
||||
let token_service: Arc<dyn TokenServicePort> = Arc::new(JwtTokenService::new(
|
||||
@@ -23,14 +23,14 @@ pub async fn create_auth_services(
|
||||
config.auth.access_token_expiry_secs,
|
||||
config.auth.refresh_token_expiry_secs,
|
||||
));
|
||||
|
||||
|
||||
// Create password hashing service
|
||||
let password_hasher = Arc::new(Argon2PasswordHasher::new());
|
||||
|
||||
|
||||
// Create PostgreSQL repositories
|
||||
let user_repository = Arc::new(UserPgRepository::new(pool.clone()));
|
||||
let session_repository = Arc::new(SessionPgRepository::new(pool.clone()));
|
||||
|
||||
|
||||
// Create authentication application service
|
||||
let mut auth_app_service = AuthApplicationService::new(
|
||||
user_repository,
|
||||
@@ -39,7 +39,7 @@ pub async fn create_auth_services(
|
||||
token_service.clone(),
|
||||
config.storage_path.clone(),
|
||||
);
|
||||
|
||||
|
||||
// Configure folder service if available
|
||||
if let Some(folder_svc) = folder_service {
|
||||
auth_app_service = auth_app_service.with_folder_service(folder_svc);
|
||||
@@ -47,9 +47,12 @@ pub async fn create_auth_services(
|
||||
|
||||
// Configure OIDC service if enabled
|
||||
if config.oidc.enabled {
|
||||
tracing::info!("Initializing OIDC service (provider: {}, issuer: {})",
|
||||
config.oidc.provider_name, config.oidc.issuer_url);
|
||||
|
||||
tracing::info!(
|
||||
"Initializing OIDC service (provider: {}, issuer: {})",
|
||||
config.oidc.provider_name,
|
||||
config.oidc.issuer_url
|
||||
);
|
||||
|
||||
let oidc_service = Arc::new(OidcService::new(config.oidc.clone()));
|
||||
auth_app_service = auth_app_service.with_oidc(oidc_service, config.oidc.clone());
|
||||
|
||||
@@ -57,12 +60,12 @@ pub async fn create_auth_services(
|
||||
tracing::warn!("Password login is DISABLED — only OIDC authentication is allowed");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Package service in Arc
|
||||
let auth_application_service = Arc::new(auth_app_service);
|
||||
|
||||
|
||||
Ok(AuthServices {
|
||||
token_service,
|
||||
auth_application_service,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+103
-66
@@ -1,19 +1,28 @@
|
||||
use sqlx::{postgres::PgPoolOptions, PgPool, Row};
|
||||
use anyhow::Result;
|
||||
use std::time::Duration;
|
||||
use crate::common::config::AppConfig;
|
||||
use anyhow::Result;
|
||||
use sqlx::{PgPool, Row, postgres::PgPoolOptions};
|
||||
use std::time::Duration;
|
||||
|
||||
pub async fn create_database_pool(config: &AppConfig) -> Result<PgPool> {
|
||||
tracing::info!("Initializing PostgreSQL connection with URL: {}",
|
||||
config.database.connection_string.replace("postgres://", "postgres://[user]:[pass]@"));
|
||||
|
||||
tracing::info!(
|
||||
"Initializing PostgreSQL connection with URL: {}",
|
||||
config
|
||||
.database
|
||||
.connection_string
|
||||
.replace("postgres://", "postgres://[user]:[pass]@")
|
||||
);
|
||||
|
||||
let mut attempt = 0;
|
||||
const MAX_ATTEMPTS: usize = 5;
|
||||
|
||||
|
||||
while attempt < MAX_ATTEMPTS {
|
||||
attempt += 1;
|
||||
tracing::info!("PostgreSQL connection attempt #{}/{}", attempt, MAX_ATTEMPTS);
|
||||
|
||||
tracing::info!(
|
||||
"PostgreSQL connection attempt #{}/{}",
|
||||
attempt,
|
||||
MAX_ATTEMPTS
|
||||
);
|
||||
|
||||
match PgPoolOptions::new()
|
||||
.max_connections(config.database.max_connections)
|
||||
.min_connections(config.database.min_connections)
|
||||
@@ -21,60 +30,76 @@ pub async fn create_database_pool(config: &AppConfig) -> Result<PgPool> {
|
||||
.idle_timeout(Duration::from_secs(config.database.idle_timeout_secs))
|
||||
.max_lifetime(Duration::from_secs(config.database.max_lifetime_secs))
|
||||
.connect(&config.database.connection_string)
|
||||
.await {
|
||||
Ok(pool) => {
|
||||
match sqlx::query("SELECT 1").execute(&pool).await {
|
||||
Ok(_) => {
|
||||
tracing::info!("PostgreSQL connection established successfully");
|
||||
|
||||
.await
|
||||
{
|
||||
Ok(pool) => {
|
||||
match sqlx::query("SELECT 1").execute(&pool).await {
|
||||
Ok(_) => {
|
||||
tracing::info!("PostgreSQL connection established successfully");
|
||||
|
||||
if !tables_exist(&pool).await {
|
||||
tracing::warn!("Database tables do not exist. Auto-applying schema...");
|
||||
if let Err(e) = apply_schema(&pool).await {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Database schema could not be applied: {}. \
|
||||
Run manually: psql -f db/schema.sql",
|
||||
e
|
||||
));
|
||||
}
|
||||
|
||||
// Verify tables were actually created
|
||||
if !tables_exist(&pool).await {
|
||||
tracing::warn!("Database tables do not exist. Auto-applying schema...");
|
||||
if let Err(e) = apply_schema(&pool).await {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Database schema could not be applied: {}. \
|
||||
Run manually: psql -f db/schema.sql", e
|
||||
));
|
||||
}
|
||||
|
||||
// Verify tables were actually created
|
||||
if !tables_exist(&pool).await {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Database schema was applied but tables still missing. \
|
||||
return Err(anyhow::anyhow!(
|
||||
"Database schema was applied but tables still missing. \
|
||||
Check db/schema.sql for errors."
|
||||
));
|
||||
}
|
||||
tracing::info!("Database schema applied and verified successfully");
|
||||
}
|
||||
|
||||
return Ok(pool);
|
||||
},
|
||||
Err(e) => {
|
||||
tracing::error!("Error verifying connection: {}", e);
|
||||
if attempt >= MAX_ATTEMPTS {
|
||||
return Err(anyhow::anyhow!("Error verifying PostgreSQL connection: {}", e));
|
||||
));
|
||||
}
|
||||
tracing::info!("Database schema applied and verified successfully");
|
||||
}
|
||||
|
||||
return Ok(pool);
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("Error verifying connection: {}", e);
|
||||
if attempt >= MAX_ATTEMPTS {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Error verifying PostgreSQL connection: {}",
|
||||
e
|
||||
));
|
||||
}
|
||||
}
|
||||
},
|
||||
Err(e) => {
|
||||
tracing::error!("Error connecting to PostgreSQL (attempt {}/{}): {}", attempt, MAX_ATTEMPTS, e);
|
||||
if attempt >= MAX_ATTEMPTS {
|
||||
return Err(anyhow::anyhow!("Error in PostgreSQL connection: {}", e));
|
||||
}
|
||||
tokio::time::sleep(Duration::from_secs(2)).await;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(
|
||||
"Error connecting to PostgreSQL (attempt {}/{}): {}",
|
||||
attempt,
|
||||
MAX_ATTEMPTS,
|
||||
e
|
||||
);
|
||||
if attempt >= MAX_ATTEMPTS {
|
||||
return Err(anyhow::anyhow!("Error in PostgreSQL connection: {}", e));
|
||||
}
|
||||
tokio::time::sleep(Duration::from_secs(2)).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Err(anyhow::anyhow!("Could not establish PostgreSQL connection after {} attempts", MAX_ATTEMPTS))
|
||||
|
||||
Err(anyhow::anyhow!(
|
||||
"Could not establish PostgreSQL connection after {} attempts",
|
||||
MAX_ATTEMPTS
|
||||
))
|
||||
}
|
||||
|
||||
/// Check whether the core auth tables exist in the database.
|
||||
async fn tables_exist(pool: &PgPool) -> bool {
|
||||
sqlx::query("SELECT EXISTS (SELECT 1 FROM pg_tables WHERE schemaname = 'auth' AND tablename = 'users')")
|
||||
.fetch_one(pool)
|
||||
.await.map(|row| row.get::<bool, _>(0))
|
||||
.unwrap_or(false)
|
||||
sqlx::query(
|
||||
"SELECT EXISTS (SELECT 1 FROM pg_tables WHERE schemaname = 'auth' AND tablename = 'users')",
|
||||
)
|
||||
.fetch_one(pool)
|
||||
.await
|
||||
.map(|row| row.get::<bool, _>(0))
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
/// Apply the embedded schema.sql to the database.
|
||||
@@ -82,15 +107,18 @@ async fn tables_exist(pool: &PgPool) -> bool {
|
||||
/// to splitting the SQL into individual statements and executing them one by one.
|
||||
async fn apply_schema(pool: &PgPool) -> Result<()> {
|
||||
let schema_sql = include_str!("../../db/schema.sql");
|
||||
|
||||
|
||||
// Attempt 1: raw_sql sends the entire script via the simple query protocol
|
||||
match sqlx::raw_sql(schema_sql).execute(pool).await {
|
||||
Ok(_) => return Ok(()),
|
||||
Err(e) => {
|
||||
tracing::warn!("raw_sql failed ({}), falling back to statement-by-statement execution", e);
|
||||
tracing::warn!(
|
||||
"raw_sql failed ({}), falling back to statement-by-statement execution",
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Attempt 2: split into individual statements respecting dollar-quoting
|
||||
let statements = split_sql_statements(schema_sql);
|
||||
for (i, stmt) in statements.iter().enumerate() {
|
||||
@@ -99,12 +127,21 @@ async fn apply_schema(pool: &PgPool) -> Result<()> {
|
||||
continue;
|
||||
}
|
||||
if let Err(e) = sqlx::raw_sql(trimmed).execute(pool).await {
|
||||
let preview = if trimmed.len() > 200 { &trimmed[..200] } else { trimmed };
|
||||
tracing::error!("Schema statement {} failed: {}\n--- SQL ---\n{}\n-----------", i + 1, e, preview);
|
||||
let preview = if trimmed.len() > 200 {
|
||||
&trimmed[..200]
|
||||
} else {
|
||||
trimmed
|
||||
};
|
||||
tracing::error!(
|
||||
"Schema statement {} failed: {}\n--- SQL ---\n{}\n-----------",
|
||||
i + 1,
|
||||
e,
|
||||
preview
|
||||
);
|
||||
return Err(anyhow::anyhow!("Schema statement {} failed: {}", i + 1, e));
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -119,7 +156,7 @@ fn split_sql_statements(sql: &str) -> Vec<String> {
|
||||
let chars: Vec<char> = sql.chars().collect();
|
||||
let len = chars.len();
|
||||
let mut i = 0;
|
||||
|
||||
|
||||
while i < len {
|
||||
// Line comment
|
||||
if i + 1 < len && chars[i] == '-' && chars[i + 1] == '-' {
|
||||
@@ -129,7 +166,7 @@ fn split_sql_statements(sql: &str) -> Vec<String> {
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
|
||||
// Block comment
|
||||
if i + 1 < len && chars[i] == '/' && chars[i + 1] == '*' {
|
||||
current.push(chars[i]);
|
||||
@@ -146,7 +183,7 @@ fn split_sql_statements(sql: &str) -> Vec<String> {
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
|
||||
// Single-quoted string
|
||||
if chars[i] == '\'' {
|
||||
current.push(chars[i]);
|
||||
@@ -167,7 +204,7 @@ fn split_sql_statements(sql: &str) -> Vec<String> {
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
|
||||
// Dollar-quoted string ($tag$...$tag$ or $$...$$)
|
||||
if chars[i] == '$' {
|
||||
let _start = i;
|
||||
@@ -203,7 +240,7 @@ fn split_sql_statements(sql: &str) -> Vec<String> {
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
|
||||
// Statement separator
|
||||
if chars[i] == ';' {
|
||||
current.push(';');
|
||||
@@ -215,16 +252,16 @@ fn split_sql_statements(sql: &str) -> Vec<String> {
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
|
||||
|
||||
current.push(chars[i]);
|
||||
i += 1;
|
||||
}
|
||||
|
||||
|
||||
// Trailing statement without semicolon
|
||||
let trimmed = current.trim().to_string();
|
||||
if !trimmed.is_empty() && trimmed != ";" {
|
||||
statements.push(trimmed);
|
||||
}
|
||||
|
||||
|
||||
statements
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,4 +3,3 @@ pub mod auth_factory;
|
||||
pub mod db;
|
||||
pub mod repositories;
|
||||
pub mod services;
|
||||
|
||||
|
||||
@@ -86,7 +86,9 @@ impl FileWritePort for CompositeFileRepository {
|
||||
content_type: String,
|
||||
content: Vec<u8>,
|
||||
) -> Result<File, DomainError> {
|
||||
self.write.save_file(name, folder_id, content_type, content).await
|
||||
self.write
|
||||
.save_file(name, folder_id, content_type, content)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn save_file_from_stream(
|
||||
@@ -96,7 +98,9 @@ impl FileWritePort for CompositeFileRepository {
|
||||
content_type: String,
|
||||
stream: std::pin::Pin<Box<dyn Stream<Item = Result<Bytes, std::io::Error>> + Send>>,
|
||||
) -> Result<File, DomainError> {
|
||||
self.write.save_file_from_stream(name, folder_id, content_type, stream).await
|
||||
self.write
|
||||
.save_file_from_stream(name, folder_id, content_type, stream)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn move_file(
|
||||
@@ -107,11 +111,7 @@ impl FileWritePort for CompositeFileRepository {
|
||||
self.write.move_file(file_id, target_folder_id).await
|
||||
}
|
||||
|
||||
async fn rename_file(
|
||||
&self,
|
||||
file_id: &str,
|
||||
new_name: &str,
|
||||
) -> Result<File, DomainError> {
|
||||
async fn rename_file(&self, file_id: &str, new_name: &str) -> Result<File, DomainError> {
|
||||
self.write.rename_file(file_id, new_name).await
|
||||
}
|
||||
|
||||
@@ -119,7 +119,11 @@ impl FileWritePort for CompositeFileRepository {
|
||||
self.write.delete_file(id).await
|
||||
}
|
||||
|
||||
async fn update_file_content(&self, file_id: &str, content: Vec<u8>) -> Result<(), DomainError> {
|
||||
async fn update_file_content(
|
||||
&self,
|
||||
file_id: &str,
|
||||
content: Vec<u8>,
|
||||
) -> Result<(), DomainError> {
|
||||
self.write.update_file_content(file_id, content).await
|
||||
}
|
||||
|
||||
@@ -130,14 +134,20 @@ impl FileWritePort for CompositeFileRepository {
|
||||
content_type: String,
|
||||
size: u64,
|
||||
) -> Result<(File, PathBuf), DomainError> {
|
||||
self.write.register_file_deferred(name, folder_id, content_type, size).await
|
||||
self.write
|
||||
.register_file_deferred(name, folder_id, content_type, size)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn move_to_trash(&self, file_id: &str) -> Result<(), DomainError> {
|
||||
self.write.move_to_trash(file_id).await
|
||||
}
|
||||
|
||||
async fn restore_from_trash(&self, file_id: &str, original_path: &str) -> Result<(), DomainError> {
|
||||
async fn restore_from_trash(
|
||||
&self,
|
||||
file_id: &str,
|
||||
original_path: &str,
|
||||
) -> Result<(), DomainError> {
|
||||
self.write.restore_from_trash(file_id, original_path).await
|
||||
}
|
||||
|
||||
|
||||
@@ -2,24 +2,26 @@ use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use tokio::{fs, time};
|
||||
use tokio::fs::File as TokioFile;
|
||||
use tokio_util::codec::{BytesCodec, FramedRead};
|
||||
use futures::{Stream, StreamExt};
|
||||
use bytes::Bytes;
|
||||
use tokio::task;
|
||||
use futures::{Stream, StreamExt};
|
||||
use mime_guess::from_path;
|
||||
use tokio::fs::File as TokioFile;
|
||||
use tokio::task;
|
||||
use tokio::{fs, time};
|
||||
use tokio_util::codec::{BytesCodec, FramedRead};
|
||||
|
||||
use crate::domain::entities::file::File;
|
||||
use crate::application::ports::storage_ports::FileReadPort;
|
||||
use crate::common::errors::DomainError;
|
||||
use crate::infrastructure::repositories::repository_errors::{FileRepositoryResult, FileRepositoryError};
|
||||
use crate::infrastructure::repositories::parallel_file_processor::ParallelFileProcessor;
|
||||
use crate::application::ports::cache_ports::MetadataCachePort;
|
||||
use crate::application::ports::storage_ports::FileReadPort;
|
||||
use crate::application::services::storage_mediator::StorageMediator;
|
||||
use crate::infrastructure::services::path_service::PathService;
|
||||
use crate::domain::services::path_service::StoragePath;
|
||||
use crate::common::config::AppConfig;
|
||||
use crate::common::errors::DomainError;
|
||||
use crate::domain::entities::file::File;
|
||||
use crate::domain::services::path_service::StoragePath;
|
||||
use crate::infrastructure::repositories::parallel_file_processor::ParallelFileProcessor;
|
||||
use crate::infrastructure::repositories::repository_errors::{
|
||||
FileRepositoryError, FileRepositoryResult,
|
||||
};
|
||||
use crate::infrastructure::services::path_service::PathService;
|
||||
|
||||
/// Repository implementation for file **read** operations.
|
||||
///
|
||||
@@ -63,12 +65,13 @@ impl FileFsReadRepository {
|
||||
Self {
|
||||
root_path: PathBuf::from("./storage"),
|
||||
storage_mediator: Arc::new(
|
||||
crate::application::services::storage_mediator::FileSystemStorageMediator::new_stub(),
|
||||
crate::application::services::storage_mediator::FileSystemStorageMediator::new_stub(
|
||||
),
|
||||
),
|
||||
id_mapping_service: Arc::new(crate::common::stubs::StubIdMappingPort),
|
||||
path_service: Arc::new(PathService::new(PathBuf::from("./storage"))),
|
||||
metadata_cache: Arc::new(
|
||||
crate::infrastructure::services::file_metadata_cache::FileMetadataCache::default()
|
||||
crate::infrastructure::services::file_metadata_cache::FileMetadataCache::default(),
|
||||
) as Arc<dyn MetadataCachePort>,
|
||||
config: AppConfig::default(),
|
||||
parallel_processor: None,
|
||||
@@ -81,36 +84,61 @@ impl FileFsReadRepository {
|
||||
self.path_service.resolve_path(storage_path)
|
||||
}
|
||||
|
||||
async fn get_file_metadata_raw(&self, abs_path: &PathBuf) -> FileRepositoryResult<(u64, u64, u64)> {
|
||||
async fn get_file_metadata_raw(
|
||||
&self,
|
||||
abs_path: &PathBuf,
|
||||
) -> FileRepositoryResult<(u64, u64, u64)> {
|
||||
// Cache first
|
||||
if let Some(cached) = self.metadata_cache.get_metadata(abs_path).await
|
||||
&& let (Some(s), Some(c), Some(m)) = (cached.size, cached.created_at, cached.modified_at) {
|
||||
return Ok((s, c, m));
|
||||
}
|
||||
&& let (Some(s), Some(c), Some(m)) =
|
||||
(cached.size, cached.created_at, cached.modified_at)
|
||||
{
|
||||
return Ok((s, c, m));
|
||||
}
|
||||
let metadata = time::timeout(self.config.timeouts.file_timeout(), fs::metadata(abs_path))
|
||||
.await
|
||||
.map_err(|_| FileRepositoryError::StorageError(format!("Timeout metadata: {}", abs_path.display())))?
|
||||
.map_err(|_| {
|
||||
FileRepositoryError::StorageError(format!(
|
||||
"Timeout metadata: {}",
|
||||
abs_path.display()
|
||||
))
|
||||
})?
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
let size = metadata.len();
|
||||
let created_at = metadata.created()
|
||||
.map(|t| t.duration_since(std::time::UNIX_EPOCH).unwrap_or_default().as_secs())
|
||||
let created_at = metadata
|
||||
.created()
|
||||
.map(|t| {
|
||||
t.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs()
|
||||
})
|
||||
.unwrap_or(0);
|
||||
let modified_at = metadata.modified()
|
||||
.map(|t| t.duration_since(std::time::UNIX_EPOCH).unwrap_or_default().as_secs())
|
||||
let modified_at = metadata
|
||||
.modified()
|
||||
.map(|t| {
|
||||
t.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs()
|
||||
})
|
||||
.unwrap_or(0);
|
||||
let _ = self.metadata_cache.refresh_metadata(abs_path).await;
|
||||
Ok((size, created_at, modified_at))
|
||||
}
|
||||
|
||||
async fn get_file_by_id(&self, id: &str) -> FileRepositoryResult<File> {
|
||||
let storage_path = self.id_mapping_service.get_path_by_id(id).await
|
||||
let storage_path = self
|
||||
.id_mapping_service
|
||||
.get_path_by_id(id)
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::Other(e.to_string()))?;
|
||||
let abs_path = self.resolve_storage_path(&storage_path);
|
||||
|
||||
if !abs_path.exists() || !abs_path.is_file() {
|
||||
return Err(FileRepositoryError::NotFound(
|
||||
format!("File {} not found at {}", id, storage_path.to_string()),
|
||||
));
|
||||
return Err(FileRepositoryError::NotFound(format!(
|
||||
"File {} not found at {}",
|
||||
id,
|
||||
storage_path.to_string()
|
||||
)));
|
||||
}
|
||||
|
||||
let (size, created_at, modified_at) = self.get_file_metadata_raw(&abs_path).await?;
|
||||
@@ -120,8 +148,14 @@ impl FileFsReadRepository {
|
||||
let mime_type = from_path(&abs_path).first_or_octet_stream().to_string();
|
||||
|
||||
File::with_timestamps(
|
||||
id.to_string(), name, storage_path, size, mime_type, None,
|
||||
created_at, modified_at,
|
||||
id.to_string(),
|
||||
name,
|
||||
storage_path,
|
||||
size,
|
||||
mime_type,
|
||||
None,
|
||||
created_at,
|
||||
modified_at,
|
||||
)
|
||||
.map_err(|e| FileRepositoryError::Other(e.to_string()))
|
||||
}
|
||||
@@ -153,18 +187,14 @@ impl FileReadPort for FileFsReadRepository {
|
||||
|
||||
async fn list_files(&self, folder_id: Option<&str>) -> Result<Vec<File>, DomainError> {
|
||||
let folder_storage_path = match folder_id {
|
||||
Some(id) => {
|
||||
match self.storage_mediator.get_folder_path(id).await {
|
||||
Ok(path) => {
|
||||
let lossy = path.to_string_lossy().to_string();
|
||||
let folder_name = path.file_name()
|
||||
.and_then(|f| f.to_str())
|
||||
.unwrap_or(&lossy);
|
||||
StoragePath::from_string(folder_name)
|
||||
}
|
||||
Err(_) => return Ok(Vec::new()),
|
||||
Some(id) => match self.storage_mediator.get_folder_path(id).await {
|
||||
Ok(path) => {
|
||||
let lossy = path.to_string_lossy().to_string();
|
||||
let folder_name = path.file_name().and_then(|f| f.to_str()).unwrap_or(&lossy);
|
||||
StoragePath::from_string(folder_name)
|
||||
}
|
||||
}
|
||||
Err(_) => return Ok(Vec::new()),
|
||||
},
|
||||
None => StoragePath::root(),
|
||||
};
|
||||
|
||||
@@ -174,16 +204,24 @@ impl FileReadPort for FileFsReadRepository {
|
||||
}
|
||||
|
||||
let mut files_result = Vec::new();
|
||||
let mut entries = fs::read_dir(&abs_folder_path).await
|
||||
let mut entries = fs::read_dir(&abs_folder_path)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
|
||||
while let Some(entry) = entries.next_entry().await
|
||||
while let Some(entry) = entries
|
||||
.next_entry()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?
|
||||
{
|
||||
let path = entry.path();
|
||||
if !path.is_file() { continue; }
|
||||
if !path.is_file() {
|
||||
continue;
|
||||
}
|
||||
let file_name = entry.file_name().to_string_lossy().to_string();
|
||||
if file_name.starts_with('.') || file_name == "folder_ids.json" || file_name == "file_ids.json" {
|
||||
if file_name.starts_with('.')
|
||||
|| file_name == "folder_ids.json"
|
||||
|| file_name == "file_ids.json"
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let metadata = match fs::metadata(&path).await {
|
||||
@@ -191,20 +229,43 @@ impl FileReadPort for FileFsReadRepository {
|
||||
Err(_) => continue,
|
||||
};
|
||||
let file_storage_path = folder_storage_path.join(&file_name);
|
||||
let id = match self.id_mapping_service.get_or_create_id(&file_storage_path).await {
|
||||
let id = match self
|
||||
.id_mapping_service
|
||||
.get_or_create_id(&file_storage_path)
|
||||
.await
|
||||
{
|
||||
Ok(id) => id,
|
||||
Err(_) => continue,
|
||||
};
|
||||
let size = metadata.len();
|
||||
let created_at = metadata.created()
|
||||
.map(|t| t.duration_since(std::time::UNIX_EPOCH).unwrap_or_default().as_secs())
|
||||
let created_at = metadata
|
||||
.created()
|
||||
.map(|t| {
|
||||
t.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs()
|
||||
})
|
||||
.unwrap_or(0);
|
||||
let modified_at = metadata.modified()
|
||||
.map(|t| t.duration_since(std::time::UNIX_EPOCH).unwrap_or_default().as_secs())
|
||||
let modified_at = metadata
|
||||
.modified()
|
||||
.map(|t| {
|
||||
t.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs()
|
||||
})
|
||||
.unwrap_or(0);
|
||||
let mime_type = from_path(&path).first_or_octet_stream().to_string();
|
||||
|
||||
match File::with_timestamps(id, file_name, file_storage_path, size, mime_type, folder_id.map(String::from), created_at, modified_at) {
|
||||
match File::with_timestamps(
|
||||
id,
|
||||
file_name,
|
||||
file_storage_path,
|
||||
size,
|
||||
mime_type,
|
||||
folder_id.map(String::from),
|
||||
created_at,
|
||||
modified_at,
|
||||
) {
|
||||
Ok(file) => files_result.push(file),
|
||||
Err(_) => continue,
|
||||
}
|
||||
@@ -216,23 +277,39 @@ impl FileReadPort for FileFsReadRepository {
|
||||
}
|
||||
|
||||
async fn get_file_content(&self, id: &str) -> Result<Vec<u8>, DomainError> {
|
||||
let file = self.get_file_by_id(id).await
|
||||
let file = self
|
||||
.get_file_by_id(id)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let abs_path = self.resolve_storage_path(file.storage_path());
|
||||
|
||||
let metadata = time::timeout(self.config.timeouts.file_timeout(), fs::metadata(&abs_path))
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", format!("Timeout metadata: {}", abs_path.display())))?
|
||||
.map_err(|_| {
|
||||
DomainError::internal_error(
|
||||
"File",
|
||||
format!("Timeout metadata: {}", abs_path.display()),
|
||||
)
|
||||
})?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let file_size = metadata.len();
|
||||
|
||||
if !self.config.resources.can_load_in_memory(file_size) {
|
||||
return Err(DomainError::internal_error("File",
|
||||
format!("File too large for memory: {} MB", file_size / (1024 * 1024))));
|
||||
return Err(DomainError::internal_error(
|
||||
"File",
|
||||
format!(
|
||||
"File too large for memory: {} MB",
|
||||
file_size / (1024 * 1024)
|
||||
),
|
||||
));
|
||||
}
|
||||
|
||||
// Parallel read for very large files
|
||||
if self.config.resources.needs_parallel_processing(file_size, &self.config.concurrency) {
|
||||
if self
|
||||
.config
|
||||
.resources
|
||||
.needs_parallel_processing(file_size, &self.config.concurrency)
|
||||
{
|
||||
let content = if let Some(processor) = &self.parallel_processor {
|
||||
processor.read_file_parallel(&abs_path).await
|
||||
} else {
|
||||
@@ -247,13 +324,14 @@ impl FileReadPort for FileFsReadRepository {
|
||||
let abs_clone = abs_path.clone();
|
||||
let chunk_size = self.config.resources.chunk_size_bytes;
|
||||
let content = task::spawn_blocking(move || -> std::io::Result<Vec<u8>> {
|
||||
use std::io::{Read, BufReader};
|
||||
use std::io::{BufReader, Read};
|
||||
let file = std::fs::File::open(&abs_clone)?;
|
||||
let mut reader = BufReader::with_capacity(chunk_size, file);
|
||||
let mut buf = Vec::with_capacity(file_size as usize);
|
||||
reader.read_to_end(&mut buf)?;
|
||||
Ok(buf)
|
||||
}).await
|
||||
})
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
return Ok(content);
|
||||
@@ -262,7 +340,12 @@ impl FileReadPort for FileFsReadRepository {
|
||||
// Small files — async read
|
||||
time::timeout(self.config.timeouts.file_timeout(), fs::read(&abs_path))
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", format!("Timeout reading: {}", abs_path.display())))?
|
||||
.map_err(|_| {
|
||||
DomainError::internal_error(
|
||||
"File",
|
||||
format!("Timeout reading: {}", abs_path.display()),
|
||||
)
|
||||
})?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))
|
||||
}
|
||||
|
||||
@@ -270,7 +353,9 @@ impl FileReadPort for FileFsReadRepository {
|
||||
&self,
|
||||
id: &str,
|
||||
) -> Result<Box<dyn Stream<Item = Result<Bytes, std::io::Error>> + Send>, DomainError> {
|
||||
let file = self.get_file_by_id(id).await
|
||||
let file = self
|
||||
.get_file_by_id(id)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let abs_path = self.resolve_storage_path(file.storage_path());
|
||||
|
||||
@@ -281,15 +366,22 @@ impl FileReadPort for FileFsReadRepository {
|
||||
let file_size = metadata.len();
|
||||
let is_large = self.config.resources.is_large_file(file_size);
|
||||
|
||||
let fh = time::timeout(self.config.timeouts.file_timeout(), TokioFile::open(&abs_path))
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout opening file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let fh = time::timeout(
|
||||
self.config.timeouts.file_timeout(),
|
||||
TokioFile::open(&abs_path),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout opening file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
|
||||
let chunk_size = if is_large { self.config.resources.chunk_size_bytes } else { 4096 };
|
||||
let chunk_size = if is_large {
|
||||
self.config.resources.chunk_size_bytes
|
||||
} else {
|
||||
4096
|
||||
};
|
||||
let codec = BytesCodec::new();
|
||||
let stream = FramedRead::with_capacity(fh, codec, chunk_size)
|
||||
.map(|r| r.map(|bm| bm.freeze()));
|
||||
let stream =
|
||||
FramedRead::with_capacity(fh, codec, chunk_size).map(|r| r.map(|bm| bm.freeze()));
|
||||
Ok(Box::new(stream))
|
||||
}
|
||||
|
||||
@@ -301,7 +393,9 @@ impl FileReadPort for FileFsReadRepository {
|
||||
) -> Result<Box<dyn Stream<Item = Result<Bytes, std::io::Error>> + Send>, DomainError> {
|
||||
use tokio::io::AsyncSeekExt;
|
||||
|
||||
let file = self.get_file_by_id(id).await
|
||||
let file = self
|
||||
.get_file_by_id(id)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let abs_path = self.resolve_storage_path(file.storage_path());
|
||||
|
||||
@@ -311,31 +405,43 @@ impl FileReadPort for FileFsReadRepository {
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let file_size = metadata.len();
|
||||
if start >= file_size {
|
||||
return Err(DomainError::internal_error("File",
|
||||
format!("Range start {} beyond file size {}", start, file_size)));
|
||||
return Err(DomainError::internal_error(
|
||||
"File",
|
||||
format!("Range start {} beyond file size {}", start, file_size),
|
||||
));
|
||||
}
|
||||
let actual_end = end.map(|e| e.min(file_size - 1)).unwrap_or(file_size - 1);
|
||||
let range_length = actual_end - start + 1;
|
||||
|
||||
let mut fh = time::timeout(self.config.timeouts.file_timeout(), TokioFile::open(&abs_path))
|
||||
let mut fh = time::timeout(
|
||||
self.config.timeouts.file_timeout(),
|
||||
TokioFile::open(&abs_path),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout opening file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.seek(std::io::SeekFrom::Start(start))
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout opening file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.seek(std::io::SeekFrom::Start(start)).await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
|
||||
let chunk_size = if range_length > 1024 * 1024 { self.config.resources.chunk_size_bytes } else { 8192 };
|
||||
let chunk_size = if range_length > 1024 * 1024 {
|
||||
self.config.resources.chunk_size_bytes
|
||||
} else {
|
||||
8192
|
||||
};
|
||||
use tokio::io::AsyncReadExt;
|
||||
let limited = fh.take(range_length);
|
||||
let codec = BytesCodec::new();
|
||||
let stream = FramedRead::with_capacity(limited, codec, chunk_size)
|
||||
.map(|r| r.map(|bm| bm.freeze()));
|
||||
let stream =
|
||||
FramedRead::with_capacity(limited, codec, chunk_size).map(|r| r.map(|bm| bm.freeze()));
|
||||
Ok(Box::new(stream))
|
||||
}
|
||||
|
||||
async fn get_file_mmap(&self, id: &str) -> Result<Bytes, DomainError> {
|
||||
use memmap2::Mmap;
|
||||
let file = self.get_file_by_id(id).await
|
||||
let file = self
|
||||
.get_file_by_id(id)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let abs_path = self.resolve_storage_path(file.storage_path());
|
||||
let path_clone = abs_path.clone();
|
||||
@@ -346,7 +452,8 @@ impl FileReadPort for FileFsReadRepository {
|
||||
let mmap = unsafe { Mmap::map(&fh) }
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
Ok(Bytes::copy_from_slice(&mmap[..]))
|
||||
}).await
|
||||
})
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?
|
||||
}
|
||||
|
||||
@@ -363,4 +470,4 @@ impl FileReadPort for FileFsReadRepository {
|
||||
_ => Ok("root".to_string()),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,25 +1,27 @@
|
||||
use async_trait::async_trait;
|
||||
use bytes::Bytes;
|
||||
use futures::{Stream, StreamExt};
|
||||
use mime_guess::from_path;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use async_trait::async_trait;
|
||||
use tokio::{fs, time};
|
||||
use tokio::fs::File as TokioFile;
|
||||
use tokio::io::AsyncWriteExt;
|
||||
use futures::{Stream, StreamExt};
|
||||
use bytes::Bytes;
|
||||
use mime_guess::from_path;
|
||||
use tokio::task;
|
||||
use tokio::{fs, time};
|
||||
|
||||
use crate::domain::entities::file::File;
|
||||
use crate::application::ports::storage_ports::FileWritePort;
|
||||
use crate::common::errors::DomainError;
|
||||
use crate::infrastructure::repositories::repository_errors::{FileRepositoryResult, FileRepositoryError};
|
||||
use crate::infrastructure::services::file_system_utils::FileSystemUtils;
|
||||
use crate::application::ports::cache_ports::MetadataCachePort;
|
||||
use crate::infrastructure::repositories::parallel_file_processor::ParallelFileProcessor;
|
||||
use crate::application::ports::storage_ports::FileWritePort;
|
||||
use crate::application::services::storage_mediator::StorageMediator;
|
||||
use crate::infrastructure::services::path_service::PathService;
|
||||
use crate::domain::services::path_service::StoragePath;
|
||||
use crate::common::config::AppConfig;
|
||||
use crate::common::errors::DomainError;
|
||||
use crate::domain::entities::file::File;
|
||||
use crate::domain::services::path_service::StoragePath;
|
||||
use crate::infrastructure::repositories::parallel_file_processor::ParallelFileProcessor;
|
||||
use crate::infrastructure::repositories::repository_errors::{
|
||||
FileRepositoryError, FileRepositoryResult,
|
||||
};
|
||||
use crate::infrastructure::services::file_system_utils::FileSystemUtils;
|
||||
use crate::infrastructure::services::path_service::PathService;
|
||||
|
||||
/// Repository implementation for file **write** operations.
|
||||
///
|
||||
@@ -47,7 +49,15 @@ impl FileFsWriteRepository {
|
||||
config: AppConfig,
|
||||
parallel_processor: Option<Arc<ParallelFileProcessor>>,
|
||||
) -> Self {
|
||||
Self { root_path, storage_mediator, id_mapping_service, path_service, metadata_cache, config, parallel_processor }
|
||||
Self {
|
||||
root_path,
|
||||
storage_mediator,
|
||||
id_mapping_service,
|
||||
path_service,
|
||||
metadata_cache,
|
||||
config,
|
||||
parallel_processor,
|
||||
}
|
||||
}
|
||||
|
||||
/// Stub for testing (does not perform real I/O).
|
||||
@@ -55,12 +65,13 @@ impl FileFsWriteRepository {
|
||||
Self {
|
||||
root_path: PathBuf::from("./storage"),
|
||||
storage_mediator: Arc::new(
|
||||
crate::application::services::storage_mediator::FileSystemStorageMediator::new_stub(),
|
||||
crate::application::services::storage_mediator::FileSystemStorageMediator::new_stub(
|
||||
),
|
||||
),
|
||||
id_mapping_service: Arc::new(crate::common::stubs::StubIdMappingPort),
|
||||
path_service: Arc::new(PathService::new(PathBuf::from("./storage"))),
|
||||
metadata_cache: Arc::new(
|
||||
crate::infrastructure::services::file_metadata_cache::FileMetadataCache::default()
|
||||
crate::infrastructure::services::file_metadata_cache::FileMetadataCache::default(),
|
||||
) as Arc<dyn MetadataCachePort>,
|
||||
config: AppConfig::default(),
|
||||
parallel_processor: None,
|
||||
@@ -78,14 +89,23 @@ impl FileFsWriteRepository {
|
||||
time::timeout(
|
||||
self.config.timeouts.dir_timeout(),
|
||||
FileSystemUtils::create_dir_with_sync(parent),
|
||||
).await
|
||||
.map_err(|_| FileRepositoryError::StorageError(format!("Timeout creating dir: {}", parent.display())))?
|
||||
)
|
||||
.await
|
||||
.map_err(|_| {
|
||||
FileRepositoryError::StorageError(format!(
|
||||
"Timeout creating dir: {}",
|
||||
parent.display()
|
||||
))
|
||||
})?
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn file_exists_at_storage_path(&self, storage_path: &StoragePath) -> FileRepositoryResult<bool> {
|
||||
async fn file_exists_at_storage_path(
|
||||
&self,
|
||||
storage_path: &StoragePath,
|
||||
) -> FileRepositoryResult<bool> {
|
||||
let abs = self.resolve_storage_path(storage_path);
|
||||
if let Some(is_file) = self.metadata_cache.is_file(&abs).await {
|
||||
return Ok(is_file);
|
||||
@@ -96,22 +116,46 @@ impl FileFsWriteRepository {
|
||||
Ok(m.is_file())
|
||||
}
|
||||
Ok(Err(_)) => Ok(false),
|
||||
Err(_) => Err(FileRepositoryError::StorageError(format!("Timeout: {}", abs.display()))),
|
||||
Err(_) => Err(FileRepositoryError::StorageError(format!(
|
||||
"Timeout: {}",
|
||||
abs.display()
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_file_metadata_raw(&self, abs_path: &PathBuf) -> FileRepositoryResult<(u64, u64, u64)> {
|
||||
async fn get_file_metadata_raw(
|
||||
&self,
|
||||
abs_path: &PathBuf,
|
||||
) -> FileRepositoryResult<(u64, u64, u64)> {
|
||||
if let Some(cached) = self.metadata_cache.get_metadata(abs_path).await
|
||||
&& let (Some(s), Some(c), Some(m)) = (cached.size, cached.created_at, cached.modified_at) {
|
||||
return Ok((s, c, m));
|
||||
}
|
||||
&& let (Some(s), Some(c), Some(m)) =
|
||||
(cached.size, cached.created_at, cached.modified_at)
|
||||
{
|
||||
return Ok((s, c, m));
|
||||
}
|
||||
let meta = time::timeout(self.config.timeouts.file_timeout(), fs::metadata(abs_path))
|
||||
.await
|
||||
.map_err(|_| FileRepositoryError::StorageError(format!("Timeout: {}", abs_path.display())))?
|
||||
.map_err(|_| {
|
||||
FileRepositoryError::StorageError(format!("Timeout: {}", abs_path.display()))
|
||||
})?
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
let s = meta.len();
|
||||
let c = meta.created().map(|t| t.duration_since(std::time::UNIX_EPOCH).unwrap_or_default().as_secs()).unwrap_or(0);
|
||||
let m = meta.modified().map(|t| t.duration_since(std::time::UNIX_EPOCH).unwrap_or_default().as_secs()).unwrap_or(0);
|
||||
let c = meta
|
||||
.created()
|
||||
.map(|t| {
|
||||
t.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs()
|
||||
})
|
||||
.unwrap_or(0);
|
||||
let m = meta
|
||||
.modified()
|
||||
.map(|t| {
|
||||
t.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs()
|
||||
})
|
||||
.unwrap_or(0);
|
||||
let _ = self.metadata_cache.refresh_metadata(abs_path).await;
|
||||
Ok((s, c, m))
|
||||
}
|
||||
@@ -159,14 +203,19 @@ impl FileFsWriteRepository {
|
||||
Err(_) => 0,
|
||||
};
|
||||
if self.config.resources.is_large_file(file_size) {
|
||||
task::spawn_blocking(move || { let _ = std::fs::remove_file(&abs_path); })
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::Other(e.to_string()))?;
|
||||
task::spawn_blocking(move || {
|
||||
let _ = std::fs::remove_file(&abs_path);
|
||||
})
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::Other(e.to_string()))?;
|
||||
} else {
|
||||
time::timeout(self.config.timeouts.file_timeout(), fs::remove_file(&abs_path))
|
||||
.await
|
||||
.map_err(|_| FileRepositoryError::StorageError("Timeout deleting file".into()))?
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
time::timeout(
|
||||
self.config.timeouts.file_timeout(),
|
||||
fs::remove_file(&abs_path),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| FileRepositoryError::StorageError("Timeout deleting file".into()))?
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -177,11 +226,14 @@ impl FileFsWriteRepository {
|
||||
match self.id_mapping_service.save_changes().await {
|
||||
Ok(_) => {
|
||||
if let Ok(verified) = self.id_mapping_service.get_path_by_id(id).await
|
||||
&& verified.to_string() == expected_path {
|
||||
return Ok(());
|
||||
}
|
||||
&& verified.to_string() == expected_path
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
if attempt == 3 {
|
||||
return Err(FileRepositoryError::Other("Failed to verify ID mapping after 3 attempts".into()));
|
||||
return Err(FileRepositoryError::Other(
|
||||
"Failed to verify ID mapping after 3 attempts".into(),
|
||||
));
|
||||
}
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||
}
|
||||
@@ -189,7 +241,12 @@ impl FileFsWriteRepository {
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||
tracing::warn!("ID mapping save retry {}: {}", attempt, e);
|
||||
}
|
||||
Err(e) => return Err(FileRepositoryError::Other(format!("Save ID mapping failed: {}", e))),
|
||||
Err(e) => {
|
||||
return Err(FileRepositoryError::Other(format!(
|
||||
"Save ID mapping failed: {}",
|
||||
e
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
@@ -229,49 +286,97 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
content: Vec<u8>,
|
||||
) -> Result<File, DomainError> {
|
||||
let folder_path = self.resolve_folder_path(&folder_id).await;
|
||||
let (file_storage_path, actual_name) = self.unique_file_path(&folder_path, &name).await.map_err(map_repo_err)?;
|
||||
let (file_storage_path, actual_name) = self
|
||||
.unique_file_path(&folder_path, &name)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
let abs_path = self.resolve_storage_path(&file_storage_path);
|
||||
self.ensure_parent_directory(&abs_path).await.map_err(map_repo_err)?;
|
||||
self.ensure_parent_directory(&abs_path)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
|
||||
let content_size = content.len() as u64;
|
||||
|
||||
// Write strategy based on file size
|
||||
if self.config.resources.needs_parallel_processing(content_size, &self.config.concurrency) {
|
||||
if self
|
||||
.config
|
||||
.resources
|
||||
.needs_parallel_processing(content_size, &self.config.concurrency)
|
||||
{
|
||||
if let Some(proc) = &self.parallel_processor {
|
||||
proc.write_file_parallel(&abs_path, &content).await.map_err(map_repo_err)?;
|
||||
proc.write_file_parallel(&abs_path, &content)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
} else {
|
||||
let proc = ParallelFileProcessor::new(self.config.clone());
|
||||
proc.write_file_parallel(&abs_path, &content).await.map_err(map_repo_err)?;
|
||||
proc.write_file_parallel(&abs_path, &content)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
}
|
||||
} else if content_size > self.config.resources.large_file_threshold_mb * 1024 * 1024 {
|
||||
let mut fh = time::timeout(self.config.timeouts.file_timeout(), TokioFile::create(&abs_path))
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout creating file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let mut fh = time::timeout(
|
||||
self.config.timeouts.file_timeout(),
|
||||
TokioFile::create(&abs_path),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout creating file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let chunk_size = self.config.resources.chunk_size_bytes;
|
||||
for chunk in content.chunks(chunk_size) {
|
||||
fh.write_all(chunk).await.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.write_all(chunk)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
}
|
||||
fh.flush().await.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
} else {
|
||||
let mut fh = time::timeout(self.config.timeouts.file_timeout(), TokioFile::create(&abs_path))
|
||||
fh.flush()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
} else {
|
||||
let mut fh = time::timeout(
|
||||
self.config.timeouts.file_timeout(),
|
||||
TokioFile::create(&abs_path),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout creating file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.write_all(&content)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.flush()
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout creating file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.write_all(&content).await.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.flush().await.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
}
|
||||
|
||||
let (size, created_at, modified_at) = self.get_file_metadata_raw(&abs_path).await.map_err(map_repo_err)?;
|
||||
let mime = if content_type.is_empty() { from_path(&abs_path).first_or_octet_stream().to_string() } else { content_type };
|
||||
let id = self.id_mapping_service.get_or_create_id(&file_storage_path).await
|
||||
let (size, created_at, modified_at) = self
|
||||
.get_file_metadata_raw(&abs_path)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
let mime = if content_type.is_empty() {
|
||||
from_path(&abs_path).first_or_octet_stream().to_string()
|
||||
} else {
|
||||
content_type
|
||||
};
|
||||
let id = self
|
||||
.id_mapping_service
|
||||
.get_or_create_id(&file_storage_path)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let path_string = file_storage_path.to_string();
|
||||
|
||||
let file = File::with_timestamps(id.clone(), actual_name, file_storage_path, size, mime, folder_id, created_at, modified_at)
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let file = File::with_timestamps(
|
||||
id.clone(),
|
||||
actual_name,
|
||||
file_storage_path,
|
||||
size,
|
||||
mime,
|
||||
folder_id,
|
||||
created_at,
|
||||
modified_at,
|
||||
)
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
|
||||
self.persist_id_mapping(&id, &path_string).await.map_err(map_repo_err)?;
|
||||
self.persist_id_mapping(&id, &path_string)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
if let Some(parent) = abs_path.parent() {
|
||||
self.metadata_cache.invalidate_directory(parent).await;
|
||||
}
|
||||
@@ -286,45 +391,86 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
mut stream: std::pin::Pin<Box<dyn Stream<Item = Result<Bytes, std::io::Error>> + Send>>,
|
||||
) -> Result<File, DomainError> {
|
||||
let folder_path = self.resolve_folder_path(&folder_id).await;
|
||||
let (file_storage_path, actual_name) = self.unique_file_path(&folder_path, &name).await.map_err(map_repo_err)?;
|
||||
let (file_storage_path, actual_name) = self
|
||||
.unique_file_path(&folder_path, &name)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
let abs_path = self.resolve_storage_path(&file_storage_path);
|
||||
self.ensure_parent_directory(&abs_path).await.map_err(map_repo_err)?;
|
||||
self.ensure_parent_directory(&abs_path)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
|
||||
let temp_path = abs_path.with_extension("tmp.upload");
|
||||
let mut fh = time::timeout(self.config.timeouts.file_timeout(), TokioFile::create(&temp_path))
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout creating temp file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let mut fh = time::timeout(
|
||||
self.config.timeouts.file_timeout(),
|
||||
TokioFile::create(&temp_path),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout creating temp file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
|
||||
let mut total_bytes: u64 = 0;
|
||||
while let Some(chunk_result) = stream.next().await {
|
||||
let chunk = chunk_result.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.write_all(&chunk).await.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let chunk =
|
||||
chunk_result.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.write_all(&chunk)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
total_bytes += chunk.len() as u64;
|
||||
}
|
||||
fh.flush().await.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.sync_all().await.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.flush()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
fh.sync_all()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
drop(fh);
|
||||
|
||||
// Atomic rename
|
||||
fs::rename(&temp_path, &abs_path).await
|
||||
fs::rename(&temp_path, &abs_path)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
|
||||
let (size, created_at, modified_at) = self.get_file_metadata_raw(&abs_path).await.map_err(map_repo_err)?;
|
||||
let mime = if content_type.is_empty() { from_path(&abs_path).first_or_octet_stream().to_string() } else { content_type };
|
||||
let id = self.id_mapping_service.get_or_create_id(&file_storage_path).await
|
||||
let (size, created_at, modified_at) = self
|
||||
.get_file_metadata_raw(&abs_path)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
let mime = if content_type.is_empty() {
|
||||
from_path(&abs_path).first_or_octet_stream().to_string()
|
||||
} else {
|
||||
content_type
|
||||
};
|
||||
let id = self
|
||||
.id_mapping_service
|
||||
.get_or_create_id(&file_storage_path)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let path_string = file_storage_path.to_string();
|
||||
let log_name = actual_name.clone();
|
||||
|
||||
let file = File::with_timestamps(id.clone(), actual_name, file_storage_path, size, mime, folder_id, created_at, modified_at)
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let file = File::with_timestamps(
|
||||
id.clone(),
|
||||
actual_name,
|
||||
file_storage_path,
|
||||
size,
|
||||
mime,
|
||||
folder_id,
|
||||
created_at,
|
||||
modified_at,
|
||||
)
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
|
||||
self.persist_id_mapping(&id, &path_string).await.map_err(map_repo_err)?;
|
||||
self.persist_id_mapping(&id, &path_string)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
if let Some(parent) = abs_path.parent() {
|
||||
self.metadata_cache.invalidate_directory(parent).await;
|
||||
}
|
||||
tracing::info!("✅ STREAMING UPLOAD COMPLETE: {} ({} bytes)", log_name, total_bytes);
|
||||
tracing::info!(
|
||||
"✅ STREAMING UPLOAD COMPLETE: {} ({} bytes)",
|
||||
log_name,
|
||||
total_bytes
|
||||
);
|
||||
Ok(file)
|
||||
}
|
||||
|
||||
@@ -339,57 +485,87 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
if !old_abs.exists() || !old_abs.is_file() {
|
||||
return Err(DomainError::not_found("File", file_id.to_string()));
|
||||
}
|
||||
let (size, created_at, modified_at) = self.get_file_metadata_raw(&old_abs).await.map_err(map_repo_err)?;
|
||||
let name = original_path.file_name()
|
||||
let (size, created_at, modified_at) = self
|
||||
.get_file_metadata_raw(&old_abs)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
let name = original_path
|
||||
.file_name()
|
||||
.ok_or_else(|| DomainError::internal_error("File", "Invalid path"))?;
|
||||
let mime = from_path(&old_abs).first_or_octet_stream().to_string();
|
||||
|
||||
// Build target path
|
||||
let target_folder_path = self.resolve_folder_path(&target_folder_id).await;
|
||||
let new_storage_path = target_folder_path.join(&name);
|
||||
if self.file_exists_at_storage_path(&new_storage_path).await.map_err(map_repo_err)? {
|
||||
return Err(DomainError::already_exists("File",
|
||||
format!("File already exists at {}", new_storage_path.to_string())));
|
||||
if self
|
||||
.file_exists_at_storage_path(&new_storage_path)
|
||||
.await
|
||||
.map_err(map_repo_err)?
|
||||
{
|
||||
return Err(DomainError::already_exists(
|
||||
"File",
|
||||
format!("File already exists at {}", new_storage_path.to_string()),
|
||||
));
|
||||
}
|
||||
let new_abs = self.resolve_storage_path(&new_storage_path);
|
||||
self.ensure_parent_directory(&new_abs).await.map_err(map_repo_err)?;
|
||||
self.ensure_parent_directory(&new_abs)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
|
||||
// Rename
|
||||
time::timeout(
|
||||
self.config.timeouts.file_timeout(),
|
||||
FileSystemUtils::rename_with_sync(&old_abs, &new_abs),
|
||||
).await
|
||||
)
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout moving file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
|
||||
// Update mapping
|
||||
self.id_mapping_service.update_path(file_id, &new_storage_path).await?;
|
||||
self.id_mapping_service
|
||||
.update_path(file_id, &new_storage_path)
|
||||
.await?;
|
||||
let _ = self.id_mapping_service.save_changes().await;
|
||||
|
||||
File::with_timestamps(file_id.to_string(), name, new_storage_path, size, mime, target_folder_id, created_at, modified_at)
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))
|
||||
File::with_timestamps(
|
||||
file_id.to_string(),
|
||||
name,
|
||||
new_storage_path,
|
||||
size,
|
||||
mime,
|
||||
target_folder_id,
|
||||
created_at,
|
||||
modified_at,
|
||||
)
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))
|
||||
}
|
||||
|
||||
async fn rename_file(
|
||||
&self,
|
||||
file_id: &str,
|
||||
new_name: &str,
|
||||
) -> Result<File, DomainError> {
|
||||
async fn rename_file(&self, file_id: &str, new_name: &str) -> Result<File, DomainError> {
|
||||
// 1. Get current file info
|
||||
let original_path = self.id_mapping_service.get_path_by_id(file_id).await?;
|
||||
let old_abs = self.resolve_storage_path(&original_path);
|
||||
if !old_abs.exists() || !old_abs.is_file() {
|
||||
return Err(DomainError::not_found("File", file_id.to_string()));
|
||||
}
|
||||
let (size, created_at, modified_at) = self.get_file_metadata_raw(&old_abs).await.map_err(map_repo_err)?;
|
||||
let (size, created_at, modified_at) = self
|
||||
.get_file_metadata_raw(&old_abs)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
|
||||
// 2. Build new path (same parent directory, different filename)
|
||||
let parent = original_path.parent()
|
||||
let parent = original_path
|
||||
.parent()
|
||||
.unwrap_or_else(|| StoragePath::new(vec![]));
|
||||
let new_storage_path = parent.join(new_name);
|
||||
if self.file_exists_at_storage_path(&new_storage_path).await.map_err(map_repo_err)? {
|
||||
return Err(DomainError::already_exists("File",
|
||||
format!("File already exists: {}", new_name)));
|
||||
if self
|
||||
.file_exists_at_storage_path(&new_storage_path)
|
||||
.await
|
||||
.map_err(map_repo_err)?
|
||||
{
|
||||
return Err(DomainError::already_exists(
|
||||
"File",
|
||||
format!("File already exists: {}", new_name),
|
||||
));
|
||||
}
|
||||
let new_abs = self.resolve_storage_path(&new_storage_path);
|
||||
let mime = from_path(&new_abs).first_or_octet_stream().to_string();
|
||||
@@ -398,12 +574,15 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
time::timeout(
|
||||
self.config.timeouts.file_timeout(),
|
||||
FileSystemUtils::rename_with_sync(&old_abs, &new_abs),
|
||||
).await
|
||||
)
|
||||
.await
|
||||
.map_err(|_| DomainError::internal_error("File", "Timeout renaming file"))?
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
|
||||
// 4. Update id→path mapping
|
||||
self.id_mapping_service.update_path(file_id, &new_storage_path).await?;
|
||||
self.id_mapping_service
|
||||
.update_path(file_id, &new_storage_path)
|
||||
.await?;
|
||||
let _ = self.id_mapping_service.save_changes().await;
|
||||
|
||||
File::with_timestamps(
|
||||
@@ -428,7 +607,9 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
self.metadata_cache.invalidate_directory(parent).await;
|
||||
}
|
||||
|
||||
self.delete_file_non_blocking(abs_path).await.map_err(map_repo_err)?;
|
||||
self.delete_file_non_blocking(abs_path)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
|
||||
// Clean up the ID mapping so we don't leave orphaned entries
|
||||
if let Err(e) = self.id_mapping_service.remove_id(id).await {
|
||||
@@ -439,7 +620,11 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn update_file_content(&self, file_id: &str, content: Vec<u8>) -> Result<(), DomainError> {
|
||||
async fn update_file_content(
|
||||
&self,
|
||||
file_id: &str,
|
||||
content: Vec<u8>,
|
||||
) -> Result<(), DomainError> {
|
||||
let storage_path = self.id_mapping_service.get_path_by_id(file_id).await?;
|
||||
let physical_path = self.resolve_storage_path(&storage_path);
|
||||
|
||||
@@ -460,9 +645,14 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
size: u64,
|
||||
) -> Result<(File, PathBuf), DomainError> {
|
||||
let folder_path = self.resolve_folder_path(&folder_id).await;
|
||||
let (file_storage_path, actual_name) = self.unique_file_path(&folder_path, &name).await.map_err(map_repo_err)?;
|
||||
let (file_storage_path, actual_name) = self
|
||||
.unique_file_path(&folder_path, &name)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
let abs_path = self.resolve_storage_path(&file_storage_path);
|
||||
self.ensure_parent_directory(&abs_path).await.map_err(map_repo_err)?;
|
||||
self.ensure_parent_directory(&abs_path)
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
|
||||
let mime = if content_type.is_empty() {
|
||||
from_path(&abs_path).first_or_octet_stream().to_string()
|
||||
@@ -474,12 +664,24 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
.unwrap_or_default()
|
||||
.as_secs();
|
||||
|
||||
let id = self.id_mapping_service.get_or_create_id(&file_storage_path).await
|
||||
let id = self
|
||||
.id_mapping_service
|
||||
.get_or_create_id(&file_storage_path)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let _ = self.id_mapping_service.save_changes().await;
|
||||
|
||||
let file = File::with_timestamps(id.clone(), actual_name, file_storage_path, size, mime, folder_id, now, now)
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
let file = File::with_timestamps(
|
||||
id.clone(),
|
||||
actual_name,
|
||||
file_storage_path,
|
||||
size,
|
||||
mime,
|
||||
folder_id,
|
||||
now,
|
||||
now,
|
||||
)
|
||||
.map_err(|e| DomainError::internal_error("File", e.to_string()))?;
|
||||
|
||||
tracing::debug!("⚡ Registered deferred file: {} -> {:?}", id, abs_path);
|
||||
Ok((file, abs_path))
|
||||
@@ -496,17 +698,21 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
|
||||
// Create trash directory
|
||||
let trash_dir = self.root_path.join(".trash").join("files");
|
||||
fs::create_dir_all(&trash_dir).await
|
||||
.map_err(|e| DomainError::internal_error("File", format!("Failed to create trash dir: {}", e)))?;
|
||||
fs::create_dir_all(&trash_dir).await.map_err(|e| {
|
||||
DomainError::internal_error("File", format!("Failed to create trash dir: {}", e))
|
||||
})?;
|
||||
|
||||
// Move file to trash
|
||||
let trash_path = trash_dir.join(file_id);
|
||||
fs::rename(&abs_path, &trash_path).await
|
||||
.map_err(|e| DomainError::internal_error("File", format!("Failed to move file to trash: {}", e)))?;
|
||||
fs::rename(&abs_path, &trash_path).await.map_err(|e| {
|
||||
DomainError::internal_error("File", format!("Failed to move file to trash: {}", e))
|
||||
})?;
|
||||
|
||||
// Update mapping to trash location
|
||||
let trash_storage_path = StoragePath::from_string(&format!(".trash/files/{}", file_id));
|
||||
self.id_mapping_service.update_path(file_id, &trash_storage_path).await?;
|
||||
self.id_mapping_service
|
||||
.update_path(file_id, &trash_storage_path)
|
||||
.await?;
|
||||
let _ = self.id_mapping_service.save_changes().await;
|
||||
|
||||
// Invalidate cache
|
||||
@@ -515,36 +721,57 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
self.metadata_cache.invalidate_directory(parent).await;
|
||||
}
|
||||
|
||||
tracing::debug!("File moved to trash: {} -> {}", file_id, trash_path.display());
|
||||
tracing::debug!(
|
||||
"File moved to trash: {} -> {}",
|
||||
file_id,
|
||||
trash_path.display()
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn restore_from_trash(&self, file_id: &str, original_path: &str) -> Result<(), DomainError> {
|
||||
async fn restore_from_trash(
|
||||
&self,
|
||||
file_id: &str,
|
||||
original_path: &str,
|
||||
) -> Result<(), DomainError> {
|
||||
// Get current path (should be in trash)
|
||||
let current_storage_path = self.id_mapping_service.get_path_by_id(file_id).await?;
|
||||
let current_abs_path = self.resolve_storage_path(¤t_storage_path);
|
||||
|
||||
if !current_abs_path.exists() {
|
||||
return Err(DomainError::not_found("File", format!("File {} not found in trash", file_id)));
|
||||
return Err(DomainError::not_found(
|
||||
"File",
|
||||
format!("File {} not found in trash", file_id),
|
||||
));
|
||||
}
|
||||
|
||||
// Ensure parent directory exists for original location
|
||||
let original_storage_path = StoragePath::from_string(original_path);
|
||||
let original_abs_path = self.resolve_storage_path(&original_storage_path);
|
||||
if let Some(parent) = original_abs_path.parent() {
|
||||
fs::create_dir_all(parent).await
|
||||
.map_err(|e| DomainError::internal_error("File", format!("Failed to create parent dir: {}", e)))?;
|
||||
fs::create_dir_all(parent).await.map_err(|e| {
|
||||
DomainError::internal_error("File", format!("Failed to create parent dir: {}", e))
|
||||
})?;
|
||||
}
|
||||
|
||||
// Move file back to original location
|
||||
fs::rename(¤t_abs_path, &original_abs_path).await
|
||||
.map_err(|e| DomainError::internal_error("File", format!("Failed to restore file: {}", e)))?;
|
||||
fs::rename(¤t_abs_path, &original_abs_path)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::internal_error("File", format!("Failed to restore file: {}", e))
|
||||
})?;
|
||||
|
||||
// Update mapping back to original path
|
||||
self.id_mapping_service.update_path(file_id, &original_storage_path).await?;
|
||||
self.id_mapping_service
|
||||
.update_path(file_id, &original_storage_path)
|
||||
.await?;
|
||||
let _ = self.id_mapping_service.save_changes().await;
|
||||
|
||||
tracing::debug!("File restored from trash: {} -> {}", file_id, original_abs_path.display());
|
||||
tracing::debug!(
|
||||
"File restored from trash: {} -> {}",
|
||||
file_id,
|
||||
original_abs_path.display()
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -555,7 +782,9 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
|
||||
// Delete the physical file if it exists
|
||||
if abs_path.exists() {
|
||||
self.delete_file_non_blocking(abs_path.clone()).await.map_err(map_repo_err)?;
|
||||
self.delete_file_non_blocking(abs_path.clone())
|
||||
.await
|
||||
.map_err(map_repo_err)?;
|
||||
}
|
||||
|
||||
// Remove ID mapping
|
||||
@@ -568,4 +797,4 @@ impl FileWritePort for FileFsWriteRepository {
|
||||
tracing::debug!("File permanently deleted: {}", file_id);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -2,8 +2,8 @@ use std::path::PathBuf;
|
||||
use tokio::fs;
|
||||
use tracing::{debug, error};
|
||||
|
||||
use crate::infrastructure::repositories::repository_errors::FolderRepositoryResult;
|
||||
use crate::infrastructure::repositories::folder_fs_repository::FolderFsRepository;
|
||||
use crate::infrastructure::repositories::repository_errors::FolderRepositoryResult;
|
||||
|
||||
// This file contains the implementation of trash-related methods
|
||||
// for the FolderFsRepository folder repository
|
||||
@@ -14,17 +14,18 @@ impl FolderFsRepository {
|
||||
fn get_trash_dir(&self) -> PathBuf {
|
||||
self.get_root_path().join(".trash").join("folders")
|
||||
}
|
||||
|
||||
|
||||
// Creates a unique path in the trash for the folder
|
||||
async fn create_trash_folder_path(&self, folder_id: &str) -> FolderRepositoryResult<PathBuf> {
|
||||
let trash_dir = self.get_trash_dir();
|
||||
|
||||
|
||||
// Ensure the trash directory exists
|
||||
if !trash_dir.exists() {
|
||||
fs::create_dir_all(&trash_dir).await
|
||||
fs::create_dir_all(&trash_dir)
|
||||
.await
|
||||
.map_err(|e| FolderRepositoryError::StorageError(e.to_string()))?;
|
||||
}
|
||||
|
||||
|
||||
// Create a unique path for the folder in the trash
|
||||
Ok(trash_dir.join(folder_id))
|
||||
}
|
||||
@@ -37,7 +38,7 @@ impl FolderFsRepository {
|
||||
/// Helper method that will be used for trash functionality
|
||||
pub(crate) async fn _trash_move_to_trash(&self, folder_id: &str) -> FolderRepositoryResult<()> {
|
||||
debug!("Moving folder to trash: {}", folder_id);
|
||||
|
||||
|
||||
// Get the physical path of the folder
|
||||
let folder_path = match self.get_mapped_folder_path(folder_id).await {
|
||||
Ok(path) => path,
|
||||
@@ -46,41 +47,55 @@ impl FolderFsRepository {
|
||||
return Err(e);
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
let folder_path_buf = PathBuf::from(folder_path.to_string());
|
||||
|
||||
|
||||
// Verify the folder exists
|
||||
if !folder_path_buf.exists() {
|
||||
return Err(FolderRepositoryError::NotFound(format!("Folder not found: {}", folder_id)));
|
||||
return Err(FolderRepositoryError::NotFound(format!(
|
||||
"Folder not found: {}",
|
||||
folder_id
|
||||
)));
|
||||
}
|
||||
|
||||
|
||||
// Create directory in the trash
|
||||
let trash_folder_path = self.create_trash_folder_path(folder_id).await?;
|
||||
|
||||
|
||||
// Physically move the folder to the trash
|
||||
match fs::rename(&folder_path_buf, &trash_folder_path).await {
|
||||
Ok(_) => {
|
||||
debug!("Folder moved to trash: {} -> {}", folder_path_buf.display(), trash_folder_path.display());
|
||||
|
||||
debug!(
|
||||
"Folder moved to trash: {} -> {}",
|
||||
folder_path_buf.display(),
|
||||
trash_folder_path.display()
|
||||
);
|
||||
|
||||
// Update the mapping to the new path in the trash
|
||||
if let Err(e) = self.update_mapped_folder_path(folder_id, &trash_folder_path).await {
|
||||
if let Err(e) = self
|
||||
.update_mapped_folder_path(folder_id, &trash_folder_path)
|
||||
.await
|
||||
{
|
||||
error!("Error updating folder mapping in trash: {}", e);
|
||||
return Err(e);
|
||||
}
|
||||
|
||||
|
||||
Ok(())
|
||||
},
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Error moving folder to trash: {}", e);
|
||||
Err(FolderRepositoryError::StorageError(e.to_string()))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Restores a folder from the trash to its original location
|
||||
pub(crate) async fn _trash_restore_from_trash(&self, folder_id: &str, original_path: &str) -> FolderRepositoryResult<()> {
|
||||
pub(crate) async fn _trash_restore_from_trash(
|
||||
&self,
|
||||
folder_id: &str,
|
||||
original_path: &str,
|
||||
) -> FolderRepositoryResult<()> {
|
||||
debug!("Restoring folder {} to {}", folder_id, original_path);
|
||||
|
||||
|
||||
// Get the current path in the trash
|
||||
let current_path = match self.get_mapped_folder_path(folder_id).await {
|
||||
Ok(path) => PathBuf::from(path),
|
||||
@@ -89,44 +104,54 @@ impl FolderFsRepository {
|
||||
return Err(e);
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
// Convert the original path to PathBuf
|
||||
let original_path_buf = PathBuf::from(original_path);
|
||||
|
||||
|
||||
// Ensure the destination parent directory exists
|
||||
if let Some(parent) = original_path_buf.parent()
|
||||
&& !parent.exists() {
|
||||
fs::create_dir_all(parent).await
|
||||
.map_err(|e| {
|
||||
error!("Error creating parent directory for restoration: {}", e);
|
||||
FolderRepositoryError::StorageError(e.to_string())
|
||||
})?;
|
||||
}
|
||||
|
||||
&& !parent.exists()
|
||||
{
|
||||
fs::create_dir_all(parent).await.map_err(|e| {
|
||||
error!("Error creating parent directory for restoration: {}", e);
|
||||
FolderRepositoryError::StorageError(e.to_string())
|
||||
})?;
|
||||
}
|
||||
|
||||
// Move the folder from the trash to its original location
|
||||
match fs::rename(¤t_path, &original_path_buf).await {
|
||||
Ok(_) => {
|
||||
debug!("Folder restored: {} -> {}", current_path.display(), original_path_buf.display());
|
||||
|
||||
debug!(
|
||||
"Folder restored: {} -> {}",
|
||||
current_path.display(),
|
||||
original_path_buf.display()
|
||||
);
|
||||
|
||||
// Update the mapping to the original path
|
||||
if let Err(e) = self.update_mapped_folder_path(folder_id, &original_path_buf).await {
|
||||
if let Err(e) = self
|
||||
.update_mapped_folder_path(folder_id, &original_path_buf)
|
||||
.await
|
||||
{
|
||||
error!("Error updating restored folder mapping: {}", e);
|
||||
return Err(e);
|
||||
}
|
||||
|
||||
|
||||
Ok(())
|
||||
},
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Error restoring folder: {}", e);
|
||||
Err(FolderRepositoryError::StorageError(e.to_string()))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Permanently deletes a folder (used by the trash)
|
||||
pub(crate) async fn _trash_delete_folder_permanently(&self, folder_id: &str) -> FolderRepositoryResult<()> {
|
||||
pub(crate) async fn _trash_delete_folder_permanently(
|
||||
&self,
|
||||
folder_id: &str,
|
||||
) -> FolderRepositoryResult<()> {
|
||||
debug!("Permanently deleting folder: {}", folder_id);
|
||||
|
||||
|
||||
// Similar to delete_folder but without additional validations
|
||||
let folder_path = match self.get_mapped_folder_path(folder_id).await {
|
||||
Ok(path) => PathBuf::from(path),
|
||||
@@ -135,13 +160,13 @@ impl FolderFsRepository {
|
||||
return Err(e);
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
// Delete the folder recursively
|
||||
if folder_path.exists() {
|
||||
match fs::remove_dir_all(&folder_path).await {
|
||||
Ok(_) => {
|
||||
debug!("Folder permanently deleted: {}", folder_path.display());
|
||||
},
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Error permanently deleting folder: {}", e);
|
||||
// Don't report error if the folder no longer exists
|
||||
@@ -151,17 +176,17 @@ impl FolderFsRepository {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Remove the mapping
|
||||
if let Err(e) = self.remove_mapped_folder_id(folder_id).await {
|
||||
error!("Error removing folder mapping: {}", e);
|
||||
return Err(e);
|
||||
}
|
||||
|
||||
|
||||
debug!("Folder permanently deleted successfully: {}", folder_id);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// Re-exports needed by the compiler
|
||||
use crate::infrastructure::repositories::repository_errors::FolderRepositoryError;
|
||||
use crate::infrastructure::repositories::repository_errors::FolderRepositoryError;
|
||||
|
||||
@@ -3,19 +3,19 @@ pub mod parallel_file_processor;
|
||||
pub mod repository_errors;
|
||||
|
||||
// Repositorios CQRS (Read/Write) + composite
|
||||
pub mod composite_file_repository;
|
||||
pub mod file_fs_read_repository;
|
||||
pub mod file_fs_write_repository;
|
||||
pub mod composite_file_repository;
|
||||
|
||||
pub mod trash_fs_repository;
|
||||
pub mod folder_fs_repository_trash;
|
||||
pub mod share_fs_repository;
|
||||
pub mod trash_fs_repository;
|
||||
|
||||
// Repositorios PostgreSQL
|
||||
pub mod pg;
|
||||
|
||||
// Re-exportar para facilitar acceso
|
||||
pub use composite_file_repository::CompositeFileRepository;
|
||||
pub use file_fs_read_repository::FileFsReadRepository;
|
||||
pub use file_fs_write_repository::FileFsWriteRepository;
|
||||
pub use composite_file_repository::CompositeFileRepository;
|
||||
pub use pg::{UserPgRepository, SessionPgRepository};
|
||||
pub use pg::{SessionPgRepository, UserPgRepository};
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use futures::future::join_all;
|
||||
use std::io::{self, SeekFrom};
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use std::io::{self, SeekFrom};
|
||||
use tokio::fs::File;
|
||||
use tokio::io::{AsyncReadExt, AsyncSeekExt, AsyncWriteExt};
|
||||
use tokio::sync::{Mutex, Semaphore};
|
||||
use tokio::task;
|
||||
use tokio::sync::{Semaphore, Mutex};
|
||||
use futures::future::join_all;
|
||||
use tracing::{info, debug, error};
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use tracing::{debug, error, info};
|
||||
|
||||
use crate::common::config::AppConfig;
|
||||
use crate::infrastructure::repositories::repository_errors::FileRepositoryError;
|
||||
@@ -39,11 +39,11 @@ impl BytesBufferPool {
|
||||
max_buffers,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Get a buffer from the pool or create a new one
|
||||
pub async fn get_buffer(&self) -> BytesMut {
|
||||
let mut buffers = self.buffers.lock().await;
|
||||
|
||||
|
||||
if let Some(mut buffer) = buffers.pop() {
|
||||
// Reuse existing buffer
|
||||
buffer.clear(); // Keep capacity, clear content
|
||||
@@ -53,14 +53,14 @@ impl BytesBufferPool {
|
||||
BytesMut::with_capacity(self.buffer_size)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Return a buffer to the pool for reuse
|
||||
pub async fn return_buffer(&self, mut buffer: BytesMut) {
|
||||
// Reset the buffer for reuse
|
||||
buffer.clear();
|
||||
|
||||
|
||||
let mut buffers = self.buffers.lock().await;
|
||||
|
||||
|
||||
// Only keep up to max_buffers
|
||||
if buffers.len() < self.max_buffers {
|
||||
buffers.push(buffer);
|
||||
@@ -85,12 +85,12 @@ impl ParallelFileProcessor {
|
||||
/// Creates a new processor instance
|
||||
pub fn new(config: AppConfig) -> Self {
|
||||
let concurrency_limiter = Arc::new(Semaphore::new(config.concurrency.max_concurrent_io));
|
||||
|
||||
|
||||
// Create BytesMut pool for efficient operations
|
||||
let chunk_size = config.resources.chunk_size_bytes;
|
||||
let max_chunks = config.concurrency.max_parallel_chunks;
|
||||
let bytes_pool = Arc::new(BytesBufferPool::new(chunk_size, max_chunks * 2));
|
||||
|
||||
|
||||
Self {
|
||||
config,
|
||||
concurrency_limiter,
|
||||
@@ -98,16 +98,16 @@ impl ParallelFileProcessor {
|
||||
bytes_pool,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Creates a new processor instance with a buffer pool
|
||||
pub fn new_with_buffer_pool(config: AppConfig, buffer_pool: Arc<BufferPool>) -> Self {
|
||||
let concurrency_limiter = Arc::new(Semaphore::new(config.concurrency.max_concurrent_io));
|
||||
|
||||
|
||||
// Create BytesMut pool for efficient operations
|
||||
let chunk_size = config.resources.chunk_size_bytes;
|
||||
let max_chunks = config.concurrency.max_parallel_chunks;
|
||||
let bytes_pool = Arc::new(BytesBufferPool::new(chunk_size, max_chunks * 2));
|
||||
|
||||
|
||||
Self {
|
||||
config,
|
||||
concurrency_limiter,
|
||||
@@ -115,34 +115,39 @@ impl ParallelFileProcessor {
|
||||
bytes_pool,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Divides a file into chunks for parallel processing
|
||||
pub fn calculate_chunks(&self, file_size: u64) -> Vec<ChunkRange> {
|
||||
// Determine if the file needs parallel processing
|
||||
let needs_parallel = self.config.resources.needs_parallel_processing(
|
||||
file_size, &self.config.concurrency
|
||||
);
|
||||
|
||||
let needs_parallel = self
|
||||
.config
|
||||
.resources
|
||||
.needs_parallel_processing(file_size, &self.config.concurrency);
|
||||
|
||||
if !needs_parallel {
|
||||
// For small files, use a single chunk
|
||||
return vec![ChunkRange {
|
||||
return vec![ChunkRange {
|
||||
index: 0,
|
||||
start: 0,
|
||||
size: file_size as usize
|
||||
size: file_size as usize,
|
||||
}];
|
||||
}
|
||||
|
||||
|
||||
// Calculate optimal number of chunks
|
||||
let chunk_count = self.config.resources.calculate_optimal_chunks(
|
||||
file_size, &self.config.concurrency
|
||||
);
|
||||
|
||||
let chunk_count = self
|
||||
.config
|
||||
.resources
|
||||
.calculate_optimal_chunks(file_size, &self.config.concurrency);
|
||||
|
||||
// Calculate size of each chunk
|
||||
let chunk_size = self.config.resources.calculate_chunk_size(file_size, chunk_count);
|
||||
|
||||
let chunk_size = self
|
||||
.config
|
||||
.resources
|
||||
.calculate_chunk_size(file_size, chunk_count);
|
||||
|
||||
// Create chunk ranges
|
||||
let mut chunks = Vec::with_capacity(chunk_count);
|
||||
|
||||
|
||||
let mut start = 0;
|
||||
for i in 0..chunk_count {
|
||||
let current_chunk_size = if i == chunk_count - 1 {
|
||||
@@ -151,275 +156,328 @@ impl ParallelFileProcessor {
|
||||
} else {
|
||||
chunk_size
|
||||
};
|
||||
|
||||
|
||||
chunks.push(ChunkRange {
|
||||
index: i,
|
||||
start,
|
||||
size: current_chunk_size,
|
||||
});
|
||||
|
||||
|
||||
start += current_chunk_size as u64;
|
||||
}
|
||||
|
||||
debug!("File size: {} bytes, divided into {} chunks of ~{} bytes each",
|
||||
file_size, chunks.len(), chunk_size);
|
||||
|
||||
|
||||
debug!(
|
||||
"File size: {} bytes, divided into {} chunks of ~{} bytes each",
|
||||
file_size,
|
||||
chunks.len(),
|
||||
chunk_size
|
||||
);
|
||||
|
||||
chunks
|
||||
}
|
||||
|
||||
|
||||
/// Reads a file in parallel and returns the complete content
|
||||
/// Optimized implementation using BytesMut to reduce memory copies
|
||||
pub async fn read_file_parallel(&self, file_path: &PathBuf) -> Result<Vec<u8>, FileRepositoryError> {
|
||||
pub async fn read_file_parallel(
|
||||
&self,
|
||||
file_path: &PathBuf,
|
||||
) -> Result<Vec<u8>, FileRepositoryError> {
|
||||
// Get file size
|
||||
let metadata = tokio::fs::metadata(file_path).await
|
||||
let metadata = tokio::fs::metadata(file_path)
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
|
||||
|
||||
let file_size = metadata.len();
|
||||
|
||||
|
||||
// Check if the file is too large for memory
|
||||
if !self.config.resources.can_load_in_memory(file_size) {
|
||||
return Err(FileRepositoryError::Other(
|
||||
format!("File too large to load in memory: {} MB (max: {} MB)",
|
||||
file_size / (1024 * 1024),
|
||||
self.config.resources.max_in_memory_file_size_mb)
|
||||
));
|
||||
return Err(FileRepositoryError::Other(format!(
|
||||
"File too large to load in memory: {} MB (max: {} MB)",
|
||||
file_size / (1024 * 1024),
|
||||
self.config.resources.max_in_memory_file_size_mb
|
||||
)));
|
||||
}
|
||||
|
||||
|
||||
// Calculate chunks
|
||||
let chunks = self.calculate_chunks(file_size);
|
||||
|
||||
|
||||
if chunks.len() == 1 {
|
||||
// For a single chunk, use simple reading with buffer pool if available
|
||||
info!("Reading file with size {}MB as a single chunk", file_size / (1024 * 1024));
|
||||
|
||||
info!(
|
||||
"Reading file with size {}MB as a single chunk",
|
||||
file_size / (1024 * 1024)
|
||||
);
|
||||
|
||||
if let Some(pool) = &self.buffer_pool {
|
||||
// Use buffer from the pool for efficient reading
|
||||
debug!("Using buffer pool for single chunk read");
|
||||
let mut buffer = pool.get_buffer().await;
|
||||
|
||||
|
||||
// If the buffer is too small, revert to standard implementation
|
||||
if buffer.capacity() < file_size as usize {
|
||||
debug!("Buffer from pool too small ({}), using standard read", buffer.capacity());
|
||||
let content = tokio::fs::read(file_path).await
|
||||
debug!(
|
||||
"Buffer from pool too small ({}), using standard read",
|
||||
buffer.capacity()
|
||||
);
|
||||
let content = tokio::fs::read(file_path)
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
|
||||
|
||||
return Ok(content);
|
||||
}
|
||||
|
||||
|
||||
// Use memory buffer from the pool
|
||||
let mut file = File::open(file_path).await
|
||||
let mut file = File::open(file_path)
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
|
||||
let read_size = file.read(buffer.as_mut_slice()).await
|
||||
|
||||
let read_size = file
|
||||
.read(buffer.as_mut_slice())
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
|
||||
|
||||
buffer.set_used(read_size);
|
||||
|
||||
|
||||
// Convert to Vec<u8>
|
||||
let content = buffer.into_vec();
|
||||
return Ok(content);
|
||||
} else {
|
||||
// Standard implementation without pool
|
||||
let content = tokio::fs::read(file_path).await
|
||||
let content = tokio::fs::read(file_path)
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
|
||||
|
||||
return Ok(content);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// For multiple chunks, use parallel reading
|
||||
info!("Reading file with size {}MB in {} parallel chunks using BytesMut",
|
||||
file_size / (1024 * 1024), chunks.len());
|
||||
|
||||
info!(
|
||||
"Reading file with size {}MB in {} parallel chunks using BytesMut",
|
||||
file_size / (1024 * 1024),
|
||||
chunks.len()
|
||||
);
|
||||
|
||||
// Create final result buffer (pre-allocated)
|
||||
let mut result = BytesMut::with_capacity(file_size as usize);
|
||||
result.resize(file_size as usize, 0);
|
||||
let result_mutex = Arc::new(Mutex::new(result));
|
||||
|
||||
|
||||
// Create tasks for each chunk
|
||||
let mut tasks = Vec::with_capacity(chunks.len());
|
||||
|
||||
|
||||
// Open file once and share it
|
||||
let file = Arc::new(File::open(file_path).await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?);
|
||||
|
||||
let file = Arc::new(
|
||||
File::open(file_path)
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?,
|
||||
);
|
||||
|
||||
// Reference to BytesMut pool
|
||||
let bytes_pool = self.bytes_pool.clone();
|
||||
|
||||
|
||||
// Process chunks in parallel
|
||||
for chunk in chunks {
|
||||
let file_clone = file.clone();
|
||||
let result_clone = result_mutex.clone();
|
||||
let semaphore_clone = self.concurrency_limiter.clone();
|
||||
let bytes_pool_clone = bytes_pool.clone();
|
||||
|
||||
|
||||
// Spawn task for this chunk - no need to copy the original data
|
||||
let task = task::spawn(async move {
|
||||
// Acquire semaphore permit
|
||||
let _permit = semaphore_clone.acquire().await.unwrap();
|
||||
|
||||
|
||||
// Get a reusable buffer from the BytesMut pool
|
||||
let mut chunk_buffer = bytes_pool_clone.get_buffer().await;
|
||||
|
||||
|
||||
// Ensure it has sufficient capacity
|
||||
if chunk_buffer.capacity() < chunk.size {
|
||||
chunk_buffer = BytesMut::with_capacity(chunk.size);
|
||||
}
|
||||
// Resize to the exact size needed
|
||||
chunk_buffer.resize(chunk.size, 0);
|
||||
|
||||
|
||||
// Create a duplicate file descriptor for independent use
|
||||
let mut file_handle = file_clone.try_clone().await?;
|
||||
|
||||
|
||||
// Position and read directly into the BytesMut
|
||||
file_handle.seek(SeekFrom::Start(chunk.start)).await?;
|
||||
let bytes_read = file_handle.read_exact(&mut chunk_buffer[..chunk.size]).await?;
|
||||
|
||||
let bytes_read = file_handle
|
||||
.read_exact(&mut chunk_buffer[..chunk.size])
|
||||
.await?;
|
||||
|
||||
if bytes_read != chunk.size {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::UnexpectedEof,
|
||||
format!("Expected to read {} bytes but got {}", chunk.size, bytes_read)
|
||||
format!(
|
||||
"Expected to read {} bytes but got {}",
|
||||
chunk.size, bytes_read
|
||||
),
|
||||
));
|
||||
}
|
||||
|
||||
|
||||
// Write to final result
|
||||
let mut result_lock = result_clone.lock().await;
|
||||
let start_pos = chunk.start as usize;
|
||||
let end_pos = start_pos + chunk.size;
|
||||
|
||||
|
||||
// Use copy_from_slice to copy from BytesMut to result buffer
|
||||
result_lock[start_pos..end_pos].copy_from_slice(&chunk_buffer[..chunk.size]);
|
||||
|
||||
|
||||
// Return the buffer to the pool for reuse
|
||||
bytes_pool_clone.return_buffer(chunk_buffer).await;
|
||||
|
||||
|
||||
// Log progress
|
||||
debug!("Chunk {} processed: {} bytes from offset {}",
|
||||
chunk.index, chunk.size, chunk.start);
|
||||
|
||||
debug!(
|
||||
"Chunk {} processed: {} bytes from offset {}",
|
||||
chunk.index, chunk.size, chunk.start
|
||||
);
|
||||
|
||||
Ok::<_, io::Error>(())
|
||||
});
|
||||
|
||||
|
||||
tasks.push(task);
|
||||
}
|
||||
|
||||
|
||||
// Wait for all tasks to complete
|
||||
let results = join_all(tasks).await;
|
||||
|
||||
|
||||
// Check for errors
|
||||
for (i, task_result) in results.into_iter().enumerate() {
|
||||
match task_result {
|
||||
Ok(Ok(())) => {},
|
||||
Ok(Ok(())) => {}
|
||||
Ok(Err(e)) => {
|
||||
error!("Error in chunk {}: {}", i, e);
|
||||
return Err(FileRepositoryError::StorageError(e.to_string()));
|
||||
},
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Task error in chunk {}: {}", i, e);
|
||||
return Err(FileRepositoryError::Other(format!("Task error: {}", e)));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Get the final result and convert to Vec<u8>
|
||||
let result_buffer = result_mutex.lock().await;
|
||||
let result_vec = result_buffer.to_vec();
|
||||
|
||||
info!("Successfully read file of {}MB in parallel with optimized BytesMut", file_size / (1024 * 1024));
|
||||
|
||||
info!(
|
||||
"Successfully read file of {}MB in parallel with optimized BytesMut",
|
||||
file_size / (1024 * 1024)
|
||||
);
|
||||
Ok(result_vec)
|
||||
}
|
||||
|
||||
|
||||
/// Writes a file in parallel from a buffer
|
||||
/// Optimized implementation using BytesMut/Bytes to reduce memory copies
|
||||
pub async fn write_file_parallel(
|
||||
&self,
|
||||
file_path: &PathBuf,
|
||||
content: &[u8]
|
||||
&self,
|
||||
file_path: &PathBuf,
|
||||
content: &[u8],
|
||||
) -> Result<(), FileRepositoryError> {
|
||||
let file_size = content.len() as u64;
|
||||
|
||||
|
||||
// Calculate chunks
|
||||
let chunks = self.calculate_chunks(file_size);
|
||||
|
||||
|
||||
if chunks.len() == 1 {
|
||||
// For a single chunk, use simple writing
|
||||
info!("Writing file with size {}MB as a single chunk", file_size / (1024 * 1024));
|
||||
|
||||
info!(
|
||||
"Writing file with size {}MB as a single chunk",
|
||||
file_size / (1024 * 1024)
|
||||
);
|
||||
|
||||
// Standard implementation (buffer pooling offers no advantages for simple writing)
|
||||
tokio::fs::write(file_path, content).await
|
||||
tokio::fs::write(file_path, content)
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
|
||||
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
|
||||
// For multiple chunks, use parallel writing
|
||||
info!("Writing file with size {}MB in {} parallel chunks using Bytes",
|
||||
file_size / (1024 * 1024), chunks.len());
|
||||
|
||||
info!(
|
||||
"Writing file with size {}MB in {} parallel chunks using Bytes",
|
||||
file_size / (1024 * 1024),
|
||||
chunks.len()
|
||||
);
|
||||
|
||||
// Create file (we don't use Mutex to reduce contention)
|
||||
let file = File::create(file_path).await
|
||||
let file = File::create(file_path)
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
|
||||
|
||||
// Convert content to Bytes (single copy step)
|
||||
let content_bytes = Bytes::copy_from_slice(content);
|
||||
|
||||
|
||||
// Create tasks for each chunk
|
||||
let mut tasks = Vec::with_capacity(chunks.len());
|
||||
|
||||
|
||||
// Process chunks in parallel
|
||||
for chunk in chunks {
|
||||
let file_clone = file.try_clone().await
|
||||
let file_clone = file
|
||||
.try_clone()
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
let semaphore_clone = self.concurrency_limiter.clone();
|
||||
|
||||
|
||||
// Create Bytes slice (doesn't copy data, only references)
|
||||
let start_idx = chunk.start as usize;
|
||||
let end_idx = start_idx + chunk.size;
|
||||
let chunk_data = content_bytes.slice(start_idx..end_idx);
|
||||
|
||||
|
||||
// Create and launch task
|
||||
let task = task::spawn(async move {
|
||||
// Acquire semaphore permit
|
||||
let _permit = semaphore_clone.acquire().await.unwrap();
|
||||
|
||||
|
||||
// Position and write
|
||||
let mut file_handle = file_clone;
|
||||
file_handle.seek(SeekFrom::Start(chunk.start)).await?;
|
||||
file_handle.write_all(&chunk_data).await?;
|
||||
|
||||
|
||||
// Log progress
|
||||
debug!("Chunk {} written: {} bytes at offset {}",
|
||||
chunk.index, chunk.size, chunk.start);
|
||||
|
||||
debug!(
|
||||
"Chunk {} written: {} bytes at offset {}",
|
||||
chunk.index, chunk.size, chunk.start
|
||||
);
|
||||
|
||||
Ok::<_, io::Error>(())
|
||||
});
|
||||
|
||||
|
||||
tasks.push(task);
|
||||
}
|
||||
|
||||
|
||||
// Wait for all tasks to complete
|
||||
let results = join_all(tasks).await;
|
||||
|
||||
|
||||
// Check for errors
|
||||
for (i, task_result) in results.into_iter().enumerate() {
|
||||
match task_result {
|
||||
Ok(Ok(())) => {},
|
||||
Ok(Ok(())) => {}
|
||||
Ok(Err(e)) => {
|
||||
error!("Error in chunk {}: {}", i, e);
|
||||
return Err(FileRepositoryError::StorageError(e.to_string()));
|
||||
},
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Task error in chunk {}: {}", i, e);
|
||||
return Err(FileRepositoryError::Other(format!("Task error: {}", e)));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Ensure everything has been written correctly
|
||||
let mut file_handle = file;
|
||||
file_handle.flush().await.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
|
||||
info!("Successfully wrote file of {}MB in parallel with optimized Bytes", file_size / (1024 * 1024));
|
||||
file_handle
|
||||
.flush()
|
||||
.await
|
||||
.map_err(|e| FileRepositoryError::StorageError(e.to_string()))?;
|
||||
|
||||
info!(
|
||||
"Successfully wrote file of {}MB in parallel with optimized Bytes",
|
||||
file_size / (1024 * 1024)
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -429,59 +487,62 @@ mod tests {
|
||||
use super::*;
|
||||
use bytes::BufMut;
|
||||
use tempfile::tempdir;
|
||||
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parallel_read_write() {
|
||||
// Create configuration with low threshold for testing
|
||||
let mut config = AppConfig::default();
|
||||
config.concurrency.min_size_for_parallel_chunks_mb = 1; // 1MB for testing
|
||||
config.concurrency.max_parallel_chunks = 4;
|
||||
|
||||
|
||||
let processor = ParallelFileProcessor::new(config);
|
||||
|
||||
|
||||
// Create temporary directory
|
||||
let temp_dir = tempdir().unwrap();
|
||||
let file_path = temp_dir.path().join("test_file.bin");
|
||||
|
||||
|
||||
// Create test data (2MB)
|
||||
let size = 2 * 1024 * 1024;
|
||||
let mut test_data = Vec::with_capacity(size);
|
||||
for i in 0..size {
|
||||
test_data.push((i % 256) as u8);
|
||||
}
|
||||
|
||||
|
||||
// Write file in parallel
|
||||
processor.write_file_parallel(&file_path, &test_data).await.unwrap();
|
||||
|
||||
processor
|
||||
.write_file_parallel(&file_path, &test_data)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Read file in parallel
|
||||
let read_data = processor.read_file_parallel(&file_path).await.unwrap();
|
||||
|
||||
|
||||
// Verify that the data is identical
|
||||
assert_eq!(test_data.len(), read_data.len());
|
||||
assert_eq!(test_data, read_data);
|
||||
}
|
||||
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_bytesmut_pool() {
|
||||
// Create pool
|
||||
let pool = BytesBufferPool::new(1024, 5);
|
||||
|
||||
|
||||
// Get buffer
|
||||
let mut buffer1 = pool.get_buffer().await;
|
||||
buffer1.put_slice(b"test data");
|
||||
assert_eq!(&buffer1[..9], b"test data");
|
||||
|
||||
|
||||
// Return buffer to the pool
|
||||
pool.return_buffer(buffer1).await;
|
||||
|
||||
|
||||
// Get another buffer (should be the same one)
|
||||
let buffer2 = pool.get_buffer().await;
|
||||
assert_eq!(buffer2.capacity(), 1024);
|
||||
|
||||
|
||||
// The buffer should be empty (cleared)
|
||||
assert_eq!(buffer2.len(), 0);
|
||||
}
|
||||
|
||||
|
||||
#[test]
|
||||
fn test_chunk_calculation() {
|
||||
// Create test configuration
|
||||
@@ -489,22 +550,22 @@ mod tests {
|
||||
config.concurrency.min_size_for_parallel_chunks_mb = 100; // 100MB
|
||||
config.concurrency.max_parallel_chunks = 4;
|
||||
config.concurrency.parallel_chunk_size_bytes = 50 * 1024 * 1024; // 50MB
|
||||
|
||||
|
||||
let processor = ParallelFileProcessor::new(config);
|
||||
|
||||
|
||||
// Small file (10MB)
|
||||
let small_file_size = 10 * 1024 * 1024;
|
||||
let chunks = processor.calculate_chunks(small_file_size);
|
||||
assert_eq!(chunks.len(), 1);
|
||||
assert_eq!(chunks[0].size as u64, small_file_size);
|
||||
|
||||
|
||||
// Large file (300MB)
|
||||
let large_file_size = 300 * 1024 * 1024;
|
||||
let chunks = processor.calculate_chunks(large_file_size);
|
||||
assert_eq!(chunks.len(), 4); // Limited to max_parallel_chunks
|
||||
|
||||
|
||||
// Verify that all chunks add up to the total size
|
||||
let total_size: u64 = chunks.iter().map(|c| c.size as u64).sum();
|
||||
assert_eq!(total_size, large_file_size);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,9 +3,11 @@ use chrono::Utc;
|
||||
use sqlx::{PgPool, Row, types::Uuid};
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::domain::entities::contact::AddressBook;
|
||||
use crate::domain::repositories::address_book_repository::{AddressBookRepository, AddressBookRepositoryResult};
|
||||
use crate::common::errors::DomainError;
|
||||
use crate::domain::entities::contact::AddressBook;
|
||||
use crate::domain::repositories::address_book_repository::{
|
||||
AddressBookRepository, AddressBookRepositoryResult,
|
||||
};
|
||||
|
||||
pub struct AddressBookPgRepository {
|
||||
pool: Arc<PgPool>,
|
||||
@@ -19,7 +21,10 @@ impl AddressBookPgRepository {
|
||||
|
||||
#[async_trait]
|
||||
impl AddressBookRepository for AddressBookPgRepository {
|
||||
async fn create_address_book(&self, address_book: AddressBook) -> AddressBookRepositoryResult<AddressBook> {
|
||||
async fn create_address_book(
|
||||
&self,
|
||||
address_book: AddressBook,
|
||||
) -> AddressBookRepositoryResult<AddressBook> {
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
INSERT INTO carddav.address_books (id, name, owner_id, description, color, is_public, created_at, updated_at)
|
||||
@@ -51,7 +56,10 @@ impl AddressBookRepository for AddressBookPgRepository {
|
||||
))
|
||||
}
|
||||
|
||||
async fn update_address_book(&self, address_book: AddressBook) -> AddressBookRepositoryResult<AddressBook> {
|
||||
async fn update_address_book(
|
||||
&self,
|
||||
address_book: AddressBook,
|
||||
) -> AddressBookRepositoryResult<AddressBook> {
|
||||
let now = Utc::now();
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
@@ -59,7 +67,7 @@ impl AddressBookRepository for AddressBookPgRepository {
|
||||
SET name = $1, description = $2, color = $3, is_public = $4, updated_at = $5
|
||||
WHERE id = $6
|
||||
RETURNING id, name, owner_id, description, color, is_public, created_at, updated_at
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(address_book.name())
|
||||
.bind(address_book.description())
|
||||
@@ -69,7 +77,9 @@ impl AddressBookRepository for AddressBookPgRepository {
|
||||
.bind(address_book.id())
|
||||
.fetch_one(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to update address book: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to update address book: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(AddressBook::from_raw(
|
||||
row.get("id"),
|
||||
@@ -88,59 +98,38 @@ impl AddressBookRepository for AddressBookPgRepository {
|
||||
r#"
|
||||
DELETE FROM carddav.address_books
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(id)
|
||||
.execute(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to delete address book: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to delete address book: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get_address_book_by_id(&self, id: &Uuid) -> AddressBookRepositoryResult<Option<AddressBook>> {
|
||||
async fn get_address_book_by_id(
|
||||
&self,
|
||||
id: &Uuid,
|
||||
) -> AddressBookRepositoryResult<Option<AddressBook>> {
|
||||
let maybe_row = sqlx::query(
|
||||
r#"
|
||||
SELECT id, name, owner_id, description, color, is_public, created_at, updated_at
|
||||
FROM carddav.address_books
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_optional(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get address book by id: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get address book by id: {}", e))
|
||||
})?;
|
||||
|
||||
let result = maybe_row.map(|row| AddressBook::from_raw(
|
||||
row.get("id"),
|
||||
row.get("name"),
|
||||
row.get("owner_id"),
|
||||
row.get("description"),
|
||||
row.get("color"),
|
||||
row.get("is_public"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
));
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
async fn get_address_books_by_owner(&self, owner_id: &str) -> AddressBookRepositoryResult<Vec<AddressBook>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT id, name, owner_id, description, color, is_public, created_at, updated_at
|
||||
FROM carddav.address_books
|
||||
WHERE owner_id = $1
|
||||
ORDER BY name
|
||||
"#
|
||||
)
|
||||
.bind(owner_id)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get address books by owner: {}", e)))?;
|
||||
|
||||
let result = rows.into_iter()
|
||||
.map(|row| AddressBook::from_raw(
|
||||
let result = maybe_row.map(|row| {
|
||||
AddressBook::from_raw(
|
||||
row.get("id"),
|
||||
row.get("name"),
|
||||
row.get("owner_id"),
|
||||
@@ -149,13 +138,54 @@ impl AddressBookRepository for AddressBookPgRepository {
|
||||
row.get("is_public"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
))
|
||||
)
|
||||
});
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
async fn get_address_books_by_owner(
|
||||
&self,
|
||||
owner_id: &str,
|
||||
) -> AddressBookRepositoryResult<Vec<AddressBook>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT id, name, owner_id, description, color, is_public, created_at, updated_at
|
||||
FROM carddav.address_books
|
||||
WHERE owner_id = $1
|
||||
ORDER BY name
|
||||
"#,
|
||||
)
|
||||
.bind(owner_id)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get address books by owner: {}", e))
|
||||
})?;
|
||||
|
||||
let result = rows
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
AddressBook::from_raw(
|
||||
row.get("id"),
|
||||
row.get("name"),
|
||||
row.get("owner_id"),
|
||||
row.get("description"),
|
||||
row.get("color"),
|
||||
row.get("is_public"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
async fn get_shared_address_books(&self, user_id: &str) -> AddressBookRepositoryResult<Vec<AddressBook>> {
|
||||
async fn get_shared_address_books(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> AddressBookRepositoryResult<Vec<AddressBook>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT a.id, a.name, a.owner_id, a.description, a.color, a.is_public, a.created_at, a.updated_at
|
||||
@@ -170,17 +200,20 @@ impl AddressBookRepository for AddressBookPgRepository {
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get shared address books: {}", e)))?;
|
||||
|
||||
let result = rows.into_iter()
|
||||
.map(|row| AddressBook::from_raw(
|
||||
row.get("id"),
|
||||
row.get("name"),
|
||||
row.get("owner_id"),
|
||||
row.get("description"),
|
||||
row.get("color"),
|
||||
row.get("is_public"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
))
|
||||
let result = rows
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
AddressBook::from_raw(
|
||||
row.get("id"),
|
||||
row.get("name"),
|
||||
row.get("owner_id"),
|
||||
row.get("description"),
|
||||
row.get("color"),
|
||||
row.get("is_public"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(result)
|
||||
@@ -193,35 +226,45 @@ impl AddressBookRepository for AddressBookPgRepository {
|
||||
FROM carddav.address_books
|
||||
WHERE is_public = true
|
||||
ORDER BY name
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get public address books: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get public address books: {}", e))
|
||||
})?;
|
||||
|
||||
let result = rows.into_iter()
|
||||
.map(|row| AddressBook::from_raw(
|
||||
row.get("id"),
|
||||
row.get("name"),
|
||||
row.get("owner_id"),
|
||||
row.get("description"),
|
||||
row.get("color"),
|
||||
row.get("is_public"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
))
|
||||
let result = rows
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
AddressBook::from_raw(
|
||||
row.get("id"),
|
||||
row.get("name"),
|
||||
row.get("owner_id"),
|
||||
row.get("description"),
|
||||
row.get("color"),
|
||||
row.get("is_public"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
async fn share_address_book(&self, address_book_id: &Uuid, user_id: &str, can_write: bool) -> AddressBookRepositoryResult<()> {
|
||||
async fn share_address_book(
|
||||
&self,
|
||||
address_book_id: &Uuid,
|
||||
user_id: &str,
|
||||
can_write: bool,
|
||||
) -> AddressBookRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO carddav.address_book_shares (address_book_id, user_id, can_write)
|
||||
VALUES ($1, $2, $3)
|
||||
ON CONFLICT (address_book_id, user_id) DO UPDATE SET can_write = $3
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(address_book_id)
|
||||
.bind(user_id)
|
||||
@@ -233,40 +276,52 @@ impl AddressBookRepository for AddressBookPgRepository {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn unshare_address_book(&self, address_book_id: &Uuid, user_id: &str) -> AddressBookRepositoryResult<()> {
|
||||
async fn unshare_address_book(
|
||||
&self,
|
||||
address_book_id: &Uuid,
|
||||
user_id: &str,
|
||||
) -> AddressBookRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
DELETE FROM carddav.address_book_shares
|
||||
WHERE address_book_id = $1 AND user_id = $2
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(address_book_id)
|
||||
.bind(user_id)
|
||||
.execute(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to unshare address book: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to unshare address book: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get_address_book_shares(&self, address_book_id: &Uuid) -> AddressBookRepositoryResult<Vec<(String, bool)>> {
|
||||
async fn get_address_book_shares(
|
||||
&self,
|
||||
address_book_id: &Uuid,
|
||||
) -> AddressBookRepositoryResult<Vec<(String, bool)>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT user_id, can_write
|
||||
FROM carddav.address_book_shares
|
||||
WHERE address_book_id = $1
|
||||
ORDER BY user_id
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(address_book_id)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get address book shares: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get address book shares: {}", e))
|
||||
})?;
|
||||
|
||||
let result = rows.into_iter()
|
||||
let result = rows
|
||||
.into_iter()
|
||||
.map(|row| (row.get("user_id"), row.get("can_write")))
|
||||
.collect();
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,9 +3,11 @@ use chrono::{DateTime, Utc};
|
||||
use sqlx::{PgPool, Row, types::Uuid};
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::domain::entities::calendar_event::CalendarEvent;
|
||||
use crate::domain::repositories::calendar_event_repository::{CalendarEventRepository, CalendarEventRepositoryResult};
|
||||
use crate::common::errors::DomainError;
|
||||
use crate::domain::entities::calendar_event::CalendarEvent;
|
||||
use crate::domain::repositories::calendar_event_repository::{
|
||||
CalendarEventRepository, CalendarEventRepositoryResult,
|
||||
};
|
||||
|
||||
pub struct CalendarEventPgRepository {
|
||||
pool: Arc<PgPool>,
|
||||
@@ -19,11 +21,14 @@ impl CalendarEventPgRepository {
|
||||
|
||||
#[async_trait]
|
||||
impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
async fn create_event(&self, event: CalendarEvent) -> CalendarEventRepositoryResult<CalendarEvent> {
|
||||
async fn create_event(
|
||||
&self,
|
||||
event: CalendarEvent,
|
||||
) -> CalendarEventRepositoryResult<CalendarEvent> {
|
||||
// This method would need a full implementation that builds the CalendarEvent
|
||||
// from the query result, using constructor methods
|
||||
// For this demonstration, we return the same event
|
||||
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO caldav.calendar_events (
|
||||
@@ -31,7 +36,7 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
all_day, rrule, created_at, updated_at, ical_uid, ical_data
|
||||
)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13)
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(event.id())
|
||||
.bind(event.calendar_id())
|
||||
@@ -48,15 +53,20 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
.bind(event.ical_data())
|
||||
.execute(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to create calendar event: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to create calendar event: {}", e))
|
||||
})?;
|
||||
|
||||
// We return the same event instead of a result
|
||||
Ok(event)
|
||||
}
|
||||
|
||||
async fn update_event(&self, event: CalendarEvent) -> CalendarEventRepositoryResult<CalendarEvent> {
|
||||
async fn update_event(
|
||||
&self,
|
||||
event: CalendarEvent,
|
||||
) -> CalendarEventRepositoryResult<CalendarEvent> {
|
||||
let now = Utc::now();
|
||||
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE caldav.calendar_events
|
||||
@@ -70,7 +80,7 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
ical_data = $8,
|
||||
updated_at = $9
|
||||
WHERE id = $10
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(event.summary())
|
||||
.bind(event.description())
|
||||
@@ -84,7 +94,9 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
.bind(event.id())
|
||||
.execute(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to update calendar event: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to update calendar event: {}", e))
|
||||
})?;
|
||||
|
||||
// In a full implementation, we would retrieve the updated event
|
||||
// For simplicity, we return the same event we received
|
||||
@@ -96,21 +108,23 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
r#"
|
||||
DELETE FROM caldav.calendar_events
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(id)
|
||||
.execute(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to delete calendar event: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to delete calendar event: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get_events_in_time_range(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
start: &DateTime<Utc>,
|
||||
end: &DateTime<Utc>
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
start: &DateTime<Utc>,
|
||||
end: &DateTime<Utc>,
|
||||
) -> CalendarEventRepositoryResult<Vec<CalendarEvent>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
@@ -127,14 +141,16 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
(rrule IS NOT NULL AND end_time >= $2)
|
||||
)
|
||||
ORDER BY start_time
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(start)
|
||||
.bind(end)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get events in time range: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get events in time range: {}", e))
|
||||
})?;
|
||||
|
||||
let mut events = Vec::new();
|
||||
for row in rows {
|
||||
@@ -152,10 +168,13 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
row.get("ical_data"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Error creating calendar event: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Error creating calendar event: {}", e))
|
||||
})?;
|
||||
events.push(event);
|
||||
}
|
||||
|
||||
|
||||
Ok(events)
|
||||
}
|
||||
|
||||
@@ -168,18 +187,20 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
created_at, updated_at, ical_uid, ical_data
|
||||
FROM caldav.calendar_events
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_optional(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get calendar event by id: {}", e)))?
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get calendar event by id: {}", e))
|
||||
})?
|
||||
.ok_or_else(|| DomainError::not_found("Calendar Event", id.to_string()))?;
|
||||
|
||||
// In a real implementation, we would build a complete CalendarEvent object
|
||||
// For simplicity, we create an object with default values to
|
||||
// demonstrate the approach without macros
|
||||
|
||||
|
||||
let event = CalendarEvent::with_id(
|
||||
row.get("id"),
|
||||
row.get("calendar_id"),
|
||||
@@ -193,13 +214,19 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
row.get("ical_uid"),
|
||||
row.get("ical_data"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at")
|
||||
).map_err(|e| DomainError::database_error(format!("Error creating calendar event: {}", e)))?;
|
||||
|
||||
row.get("updated_at"),
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Error creating calendar event: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(event)
|
||||
}
|
||||
|
||||
async fn list_events_by_calendar(&self, calendar_id: &Uuid) -> CalendarEventRepositoryResult<Vec<CalendarEvent>> {
|
||||
|
||||
async fn list_events_by_calendar(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
) -> CalendarEventRepositoryResult<Vec<CalendarEvent>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -209,12 +236,14 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
FROM caldav.calendar_events
|
||||
WHERE calendar_id = $1
|
||||
ORDER BY start_time
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get events by calendar: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get events by calendar: {}", e))
|
||||
})?;
|
||||
|
||||
let mut events = Vec::new();
|
||||
for row in rows {
|
||||
@@ -232,16 +261,23 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
row.get("ical_data"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Error creating calendar event: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Error creating calendar event: {}", e))
|
||||
})?;
|
||||
events.push(event);
|
||||
}
|
||||
|
||||
|
||||
Ok(events)
|
||||
}
|
||||
|
||||
async fn find_events_by_summary(&self, calendar_id: &Uuid, summary: &str) -> CalendarEventRepositoryResult<Vec<CalendarEvent>> {
|
||||
|
||||
async fn find_events_by_summary(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
summary: &str,
|
||||
) -> CalendarEventRepositoryResult<Vec<CalendarEvent>> {
|
||||
let search_pattern = format!("%{}%", summary);
|
||||
|
||||
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -251,13 +287,15 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
FROM caldav.calendar_events
|
||||
WHERE calendar_id = $1 AND summary ILIKE $2
|
||||
ORDER BY start_time
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(&search_pattern)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to find events by summary: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to find events by summary: {}", e))
|
||||
})?;
|
||||
|
||||
let mut events = Vec::new();
|
||||
for row in rows {
|
||||
@@ -275,14 +313,21 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
row.get("ical_data"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Error creating calendar event: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Error creating calendar event: {}", e))
|
||||
})?;
|
||||
events.push(event);
|
||||
}
|
||||
|
||||
|
||||
Ok(events)
|
||||
}
|
||||
|
||||
async fn find_event_by_ical_uid(&self, calendar_id: &Uuid, ical_uid: &str) -> CalendarEventRepositoryResult<Option<CalendarEvent>> {
|
||||
|
||||
async fn find_event_by_ical_uid(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
ical_uid: &str,
|
||||
) -> CalendarEventRepositoryResult<Option<CalendarEvent>> {
|
||||
let row_opt = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -291,13 +336,15 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
created_at, updated_at, ical_uid, ical_data
|
||||
FROM caldav.calendar_events
|
||||
WHERE calendar_id = $1 AND ical_uid = $2
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(ical_uid)
|
||||
.fetch_optional(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get calendar event by UID: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get calendar event by UID: {}", e))
|
||||
})?;
|
||||
|
||||
match row_opt {
|
||||
Some(row) => {
|
||||
@@ -315,49 +362,62 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
row.get("ical_data"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Error creating calendar event: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Error creating calendar event: {}", e))
|
||||
})?;
|
||||
Ok(Some(event))
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
async fn count_events_in_calendar(&self, calendar_id: &Uuid) -> CalendarEventRepositoryResult<i64> {
|
||||
|
||||
async fn count_events_in_calendar(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
) -> CalendarEventRepositoryResult<i64> {
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT COUNT(*) as count
|
||||
FROM caldav.calendar_events
|
||||
WHERE calendar_id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.fetch_one(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to count events in calendar: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to count events in calendar: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(row.get::<i64, _>("count"))
|
||||
}
|
||||
|
||||
async fn delete_all_events_in_calendar(&self, calendar_id: &Uuid) -> CalendarEventRepositoryResult<i64> {
|
||||
|
||||
async fn delete_all_events_in_calendar(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
) -> CalendarEventRepositoryResult<i64> {
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
DELETE FROM caldav.calendar_events
|
||||
WHERE calendar_id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.execute(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to delete all events in calendar: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to delete all events in calendar: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(result.rows_affected() as i64)
|
||||
}
|
||||
|
||||
|
||||
async fn list_events_by_calendar_paginated(
|
||||
&self,
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
limit: i64,
|
||||
offset: i64
|
||||
offset: i64,
|
||||
) -> CalendarEventRepositoryResult<Vec<CalendarEvent>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
@@ -369,14 +429,19 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
WHERE calendar_id = $1
|
||||
ORDER BY start_time
|
||||
LIMIT $2 OFFSET $3
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(limit)
|
||||
.bind(offset)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get paginated events by calendar: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!(
|
||||
"Failed to get paginated events by calendar: {}",
|
||||
e
|
||||
))
|
||||
})?;
|
||||
|
||||
let mut events = Vec::new();
|
||||
for row in rows {
|
||||
@@ -394,18 +459,21 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
row.get("ical_data"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Error creating calendar event: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Error creating calendar event: {}", e))
|
||||
})?;
|
||||
events.push(event);
|
||||
}
|
||||
|
||||
|
||||
Ok(events)
|
||||
}
|
||||
|
||||
|
||||
async fn find_recurring_events_in_range(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
start: &DateTime<Utc>,
|
||||
end: &DateTime<Utc>
|
||||
end: &DateTime<Utc>,
|
||||
) -> CalendarEventRepositoryResult<Vec<CalendarEvent>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
@@ -419,14 +487,16 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
AND end_time >= $2
|
||||
AND start_time <= $3
|
||||
ORDER BY start_time
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(start)
|
||||
.bind(end)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to find recurring events in range: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to find recurring events in range: {}", e))
|
||||
})?;
|
||||
|
||||
let mut events = Vec::new();
|
||||
for row in rows {
|
||||
@@ -444,10 +514,13 @@ impl CalendarEventRepository for CalendarEventPgRepository {
|
||||
row.get("ical_data"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Error creating calendar event: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Error creating calendar event: {}", e))
|
||||
})?;
|
||||
events.push(event);
|
||||
}
|
||||
|
||||
|
||||
Ok(events)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,9 +3,11 @@ use chrono::Utc;
|
||||
use sqlx::{PgPool, Row, types::Uuid};
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::domain::entities::calendar::Calendar;
|
||||
use crate::domain::repositories::calendar_repository::{CalendarRepository, CalendarRepositoryResult};
|
||||
use crate::common::errors::DomainError;
|
||||
use crate::domain::entities::calendar::Calendar;
|
||||
use crate::domain::repositories::calendar_repository::{
|
||||
CalendarRepository, CalendarRepositoryResult,
|
||||
};
|
||||
|
||||
pub struct CalendarPgRepository {
|
||||
pool: Arc<PgPool>,
|
||||
@@ -38,7 +40,7 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
.fetch_one(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to create calendar: {}", e)))?;
|
||||
|
||||
|
||||
// Build the Calendar object using its with_id constructor
|
||||
let result = Calendar::with_id(
|
||||
row.get("id"),
|
||||
@@ -48,7 +50,10 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
row.get("color"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Failed to create calendar object: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to create calendar object: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
@@ -61,7 +66,7 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
SET name = $1, description = $2, color = $3, is_public = $4, updated_at = $5
|
||||
WHERE id = $6
|
||||
RETURNING id, name, owner_id, description, color, is_public, created_at, updated_at
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar.name())
|
||||
.bind(calendar.description())
|
||||
@@ -72,7 +77,7 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
.fetch_one(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to update calendar: {}", e)))?;
|
||||
|
||||
|
||||
// Build the Calendar object using its with_id constructor
|
||||
let result = Calendar::with_id(
|
||||
row.get("id"),
|
||||
@@ -82,7 +87,10 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
row.get("color"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Failed to create calendar object: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to create calendar object: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
@@ -92,7 +100,7 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
r#"
|
||||
DELETE FROM caldav.calendars
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(id)
|
||||
.execute(&*self.pool)
|
||||
@@ -108,7 +116,7 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
SELECT id, name, owner_id, description, color, is_public, created_at, updated_at
|
||||
FROM caldav.calendars
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_optional(&*self.pool)
|
||||
@@ -124,24 +132,32 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
row.get("color"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Failed to create calendar object: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to create calendar object: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(calendar)
|
||||
}
|
||||
|
||||
async fn list_calendars_by_owner(&self, owner_id: &str) -> CalendarRepositoryResult<Vec<Calendar>> {
|
||||
async fn list_calendars_by_owner(
|
||||
&self,
|
||||
owner_id: &str,
|
||||
) -> CalendarRepositoryResult<Vec<Calendar>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT id, name, owner_id, description, color, is_public, created_at, updated_at
|
||||
FROM caldav.calendars
|
||||
WHERE owner_id = $1
|
||||
ORDER BY name
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(owner_id)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get calendars by owner: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get calendars by owner: {}", e))
|
||||
})?;
|
||||
|
||||
let mut calendars = Vec::new();
|
||||
for row in rows {
|
||||
@@ -153,27 +169,38 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
row.get("color"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Failed to create calendar object: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to create calendar object: {}", e))
|
||||
})?;
|
||||
calendars.push(calendar);
|
||||
}
|
||||
|
||||
Ok(calendars)
|
||||
}
|
||||
|
||||
async fn find_calendar_by_name_and_owner(&self, name: &str, owner_id: &str) -> CalendarRepositoryResult<Calendar> {
|
||||
async fn find_calendar_by_name_and_owner(
|
||||
&self,
|
||||
name: &str,
|
||||
owner_id: &str,
|
||||
) -> CalendarRepositoryResult<Calendar> {
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT id, name, owner_id, description, color, is_public, created_at, updated_at
|
||||
FROM caldav.calendars
|
||||
WHERE name = $1 AND owner_id = $2
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(name)
|
||||
.bind(owner_id)
|
||||
.fetch_optional(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to find calendar by name and owner: {}", e)))?
|
||||
.ok_or_else(|| DomainError::not_found("Calendar", format!("{} (owned by {})", name, owner_id)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to find calendar by name and owner: {}", e))
|
||||
})?
|
||||
.ok_or_else(|| {
|
||||
DomainError::not_found("Calendar", format!("{} (owned by {})", name, owner_id))
|
||||
})?;
|
||||
|
||||
let calendar = Calendar::with_id(
|
||||
row.get("id"),
|
||||
@@ -183,12 +210,18 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
row.get("color"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Failed to create calendar object: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to create calendar object: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(calendar)
|
||||
}
|
||||
|
||||
async fn list_calendars_shared_with_user(&self, user_id: &str) -> CalendarRepositoryResult<Vec<Calendar>> {
|
||||
async fn list_calendars_shared_with_user(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> CalendarRepositoryResult<Vec<Calendar>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT c.id, c.name, c.owner_id, c.description, c.color, c.is_public, c.created_at, c.updated_at
|
||||
@@ -213,14 +246,21 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
row.get("color"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Failed to create calendar object: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to create calendar object: {}", e))
|
||||
})?;
|
||||
calendars.push(calendar);
|
||||
}
|
||||
|
||||
Ok(calendars)
|
||||
}
|
||||
|
||||
async fn list_public_calendars(&self, limit: i64, offset: i64) -> CalendarRepositoryResult<Vec<Calendar>> {
|
||||
async fn list_public_calendars(
|
||||
&self,
|
||||
limit: i64,
|
||||
offset: i64,
|
||||
) -> CalendarRepositoryResult<Vec<Calendar>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT id, name, owner_id, description, color, is_public, created_at, updated_at
|
||||
@@ -228,13 +268,15 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
WHERE is_public = true
|
||||
ORDER BY name
|
||||
LIMIT $1 OFFSET $2
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(limit)
|
||||
.bind(offset)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get public calendars: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get public calendars: {}", e))
|
||||
})?;
|
||||
|
||||
let mut calendars = Vec::new();
|
||||
for row in rows {
|
||||
@@ -243,17 +285,24 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
row.get("name"),
|
||||
row.get("owner_id"),
|
||||
row.get("description"),
|
||||
row.get("color"),
|
||||
row.get("color"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
).map_err(|e| DomainError::database_error(format!("Failed to create calendar object: {}", e)))?;
|
||||
)
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to create calendar object: {}", e))
|
||||
})?;
|
||||
calendars.push(calendar);
|
||||
}
|
||||
|
||||
Ok(calendars)
|
||||
}
|
||||
|
||||
async fn user_has_calendar_access(&self, calendar_id: &Uuid, user_id: &str) -> CalendarRepositoryResult<bool> {
|
||||
async fn user_has_calendar_access(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
user_id: &str,
|
||||
) -> CalendarRepositoryResult<bool> {
|
||||
// Check if the user is the owner of the calendar or has a share
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
@@ -264,31 +313,39 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
SELECT 1 FROM caldav.calendar_shares s
|
||||
WHERE s.calendar_id = $1 AND s.user_id = $2
|
||||
) as has_access
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(user_id)
|
||||
.fetch_one(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to check calendar access: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to check calendar access: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(row.get::<bool, _>("has_access"))
|
||||
}
|
||||
|
||||
async fn share_calendar(&self, calendar_id: &Uuid, user_id: &str, access_level: &str) -> CalendarRepositoryResult<()> {
|
||||
async fn share_calendar(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
user_id: &str,
|
||||
access_level: &str,
|
||||
) -> CalendarRepositoryResult<()> {
|
||||
// Validate access level
|
||||
if !["read", "write", "owner"].contains(&access_level) {
|
||||
return Err(DomainError::validation_error(
|
||||
format!("Invalid access level: '{}'. Must be 'read', 'write', or 'owner'", access_level)
|
||||
));
|
||||
return Err(DomainError::validation_error(format!(
|
||||
"Invalid access level: '{}'. Must be 'read', 'write', or 'owner'",
|
||||
access_level
|
||||
)));
|
||||
}
|
||||
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO caldav.calendar_shares (calendar_id, user_id, access_level)
|
||||
VALUES ($1, $2, $3)
|
||||
ON CONFLICT (calendar_id, user_id) DO UPDATE SET access_level = $3
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(user_id)
|
||||
@@ -300,12 +357,16 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_calendar_sharing(&self, calendar_id: &Uuid, user_id: &str) -> CalendarRepositoryResult<()> {
|
||||
async fn remove_calendar_sharing(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
user_id: &str,
|
||||
) -> CalendarRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
DELETE FROM caldav.calendar_shares
|
||||
WHERE calendar_id = $1 AND user_id = $2
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(user_id)
|
||||
@@ -316,19 +377,24 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get_calendar_shares(&self, calendar_id: &Uuid) -> CalendarRepositoryResult<Vec<(String, String)>> {
|
||||
async fn get_calendar_shares(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
) -> CalendarRepositoryResult<Vec<(String, String)>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT user_id, access_level
|
||||
FROM caldav.calendar_shares
|
||||
WHERE calendar_id = $1
|
||||
ORDER BY user_id
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get calendar shares: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get calendar shares: {}", e))
|
||||
})?;
|
||||
|
||||
let mut shares = Vec::new();
|
||||
for row in rows {
|
||||
@@ -337,76 +403,100 @@ impl CalendarRepository for CalendarPgRepository {
|
||||
|
||||
Ok(shares)
|
||||
}
|
||||
|
||||
async fn get_calendar_property(&self, calendar_id: &Uuid, property_name: &str) -> CalendarRepositoryResult<Option<String>> {
|
||||
|
||||
async fn get_calendar_property(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
property_name: &str,
|
||||
) -> CalendarRepositoryResult<Option<String>> {
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT value
|
||||
FROM caldav.calendar_properties
|
||||
WHERE calendar_id = $1 AND name = $2
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(property_name)
|
||||
.fetch_optional(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get calendar property: {}", e)))?;
|
||||
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get calendar property: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(row.map(|r| r.get("value")))
|
||||
}
|
||||
|
||||
async fn set_calendar_property(&self, calendar_id: &Uuid, property_name: &str, property_value: &str) -> CalendarRepositoryResult<()> {
|
||||
|
||||
async fn set_calendar_property(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
property_name: &str,
|
||||
property_value: &str,
|
||||
) -> CalendarRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO caldav.calendar_properties (calendar_id, name, value)
|
||||
VALUES ($1, $2, $3)
|
||||
ON CONFLICT (calendar_id, name) DO UPDATE SET value = $3
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(property_name)
|
||||
.bind(property_value)
|
||||
.execute(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to set calendar property: {}", e)))?;
|
||||
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to set calendar property: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_calendar_property(&self, calendar_id: &Uuid, property_name: &str) -> CalendarRepositoryResult<()> {
|
||||
|
||||
async fn remove_calendar_property(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
property_name: &str,
|
||||
) -> CalendarRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
DELETE FROM caldav.calendar_properties
|
||||
WHERE calendar_id = $1 AND name = $2
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.bind(property_name)
|
||||
.execute(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to remove calendar property: {}", e)))?;
|
||||
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to remove calendar property: {}", e))
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get_calendar_properties(&self, calendar_id: &Uuid) -> CalendarRepositoryResult<std::collections::HashMap<String, String>> {
|
||||
|
||||
async fn get_calendar_properties(
|
||||
&self,
|
||||
calendar_id: &Uuid,
|
||||
) -> CalendarRepositoryResult<std::collections::HashMap<String, String>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT name, value
|
||||
FROM caldav.calendar_properties
|
||||
WHERE calendar_id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(calendar_id)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get calendar properties: {}", e)))?;
|
||||
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get calendar properties: {}", e))
|
||||
})?;
|
||||
|
||||
let mut properties = std::collections::HashMap::new();
|
||||
for row in rows {
|
||||
properties.insert(row.get("name"), row.get("value"));
|
||||
}
|
||||
|
||||
|
||||
Ok(properties)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,223 +1,276 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{PgPool, Row, types::Uuid};
|
||||
use std::sync::Arc;
|
||||
use chrono::Utc;
|
||||
use serde_json::Value as JsonValue;
|
||||
|
||||
use crate::domain::entities::contact::{Contact, ContactGroup};
|
||||
use crate::domain::repositories::contact_repository::{ContactGroupRepository, ContactRepositoryResult};
|
||||
use crate::common::errors::{DomainError, ErrorKind};
|
||||
use super::contact_persistence_dto::{
|
||||
emails_from_persistence, phones_from_persistence, addresses_from_persistence,
|
||||
EmailPersistenceDto, PhonePersistenceDto, AddressPersistenceDto,
|
||||
};
|
||||
|
||||
pub struct ContactGroupPgRepository {
|
||||
pool: Arc<PgPool>,
|
||||
}
|
||||
|
||||
impl ContactGroupPgRepository {
|
||||
pub fn new(pool: Arc<PgPool>) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ContactGroupRepository for ContactGroupPgRepository {
|
||||
async fn create_group(&self, group: ContactGroup) -> ContactRepositoryResult<ContactGroup> {
|
||||
sqlx::query(
|
||||
"INSERT INTO carddav.contact_groups (id, address_book_id, name, created_at, updated_at) VALUES ($1, $2, $3, $4, $5)"
|
||||
)
|
||||
.bind(group.id())
|
||||
.bind(group.address_book_id())
|
||||
.bind(group.name())
|
||||
.bind(group.created_at())
|
||||
.bind(group.updated_at())
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to create group: {}", e)))?;
|
||||
|
||||
Ok(group)
|
||||
}
|
||||
|
||||
async fn update_group(&self, group: ContactGroup) -> ContactRepositoryResult<ContactGroup> {
|
||||
sqlx::query(
|
||||
"UPDATE carddav.contact_groups SET name = $1, updated_at = $2 WHERE id = $3"
|
||||
)
|
||||
.bind(group.name())
|
||||
.bind(Utc::now())
|
||||
.bind(group.id())
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to update group: {}", e)))?;
|
||||
|
||||
Ok(group)
|
||||
}
|
||||
|
||||
async fn delete_group(&self, id: &Uuid) -> ContactRepositoryResult<()> {
|
||||
// Delete memberships first
|
||||
sqlx::query("DELETE FROM carddav.group_memberships WHERE group_id = $1")
|
||||
.bind(id)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to delete group memberships: {}", e)))?;
|
||||
|
||||
sqlx::query("DELETE FROM carddav.contact_groups WHERE id = $1")
|
||||
.bind(id)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to delete group: {}", e)))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get_group_by_id(&self, id: &Uuid) -> ContactRepositoryResult<Option<ContactGroup>> {
|
||||
let row = sqlx::query(
|
||||
"SELECT id, address_book_id, name, created_at, updated_at FROM carddav.contact_groups WHERE id = $1"
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_optional(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to get group: {}", e)))?;
|
||||
|
||||
match row {
|
||||
Some(row) => {
|
||||
let group = ContactGroup::from_raw(
|
||||
row.get::<Uuid, _>("id"),
|
||||
row.get::<Uuid, _>("address_book_id"),
|
||||
row.get::<String, _>("name"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
);
|
||||
Ok(Some(group))
|
||||
},
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_groups_by_address_book(&self, address_book_id: &Uuid) -> ContactRepositoryResult<Vec<ContactGroup>> {
|
||||
let rows = sqlx::query(
|
||||
"SELECT id, address_book_id, name, created_at, updated_at FROM carddav.contact_groups WHERE address_book_id = $1 ORDER BY name"
|
||||
)
|
||||
.bind(address_book_id)
|
||||
.fetch_all(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to list groups: {}", e)))?;
|
||||
|
||||
Ok(rows.into_iter().map(|row| {
|
||||
ContactGroup::from_raw(
|
||||
row.get::<Uuid, _>("id"),
|
||||
row.get::<Uuid, _>("address_book_id"),
|
||||
row.get::<String, _>("name"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
)
|
||||
}).collect())
|
||||
}
|
||||
|
||||
async fn add_contact_to_group(&self, group_id: &Uuid, contact_id: &Uuid) -> ContactRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
"INSERT INTO carddav.group_memberships (group_id, contact_id) VALUES ($1, $2) ON CONFLICT DO NOTHING"
|
||||
)
|
||||
.bind(group_id)
|
||||
.bind(contact_id)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to add contact to group: {}", e)))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_contact_from_group(&self, group_id: &Uuid, contact_id: &Uuid) -> ContactRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
"DELETE FROM carddav.group_memberships WHERE group_id = $1 AND contact_id = $2"
|
||||
)
|
||||
.bind(group_id)
|
||||
.bind(contact_id)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to remove contact from group: {}", e)))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get_contacts_in_group(&self, group_id: &Uuid) -> ContactRepositoryResult<Vec<Contact>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
c.id, c.address_book_id, c.uid, c.full_name, c.first_name, c.last_name, c.nickname,
|
||||
c.email, c.phone, c.address, c.organization, c.title, c.notes, c.photo_url,
|
||||
c.birthday, c.anniversary, c.vcard, c.etag, c.created_at, c.updated_at
|
||||
FROM carddav.contacts c
|
||||
INNER JOIN carddav.group_memberships gm ON c.id = gm.contact_id
|
||||
WHERE gm.group_id = $1
|
||||
ORDER BY c.full_name, c.first_name, c.last_name
|
||||
"#
|
||||
)
|
||||
.bind(group_id)
|
||||
.fetch_all(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to get contacts in group: {}", e)))?;
|
||||
|
||||
let mut contacts = Vec::new();
|
||||
for row in &rows {
|
||||
let email_json: JsonValue = row.get("email");
|
||||
let phone_json: JsonValue = row.get("phone");
|
||||
let address_json: JsonValue = row.get("address");
|
||||
|
||||
let emails = serde_json::from_value::<Vec<EmailPersistenceDto>>(email_json)
|
||||
.map(emails_from_persistence)
|
||||
.unwrap_or_default();
|
||||
let phones = serde_json::from_value::<Vec<PhonePersistenceDto>>(phone_json)
|
||||
.map(phones_from_persistence)
|
||||
.unwrap_or_default();
|
||||
let addresses = serde_json::from_value::<Vec<AddressPersistenceDto>>(address_json)
|
||||
.map(addresses_from_persistence)
|
||||
.unwrap_or_default();
|
||||
|
||||
contacts.push(Contact::from_raw(
|
||||
row.get("id"),
|
||||
row.get("address_book_id"),
|
||||
row.get("uid"),
|
||||
row.get::<Option<String>, _>("full_name"),
|
||||
row.get::<Option<String>, _>("first_name"),
|
||||
row.get::<Option<String>, _>("last_name"),
|
||||
row.get::<Option<String>, _>("nickname"),
|
||||
emails,
|
||||
phones,
|
||||
addresses,
|
||||
row.get::<Option<String>, _>("organization"),
|
||||
row.get::<Option<String>, _>("title"),
|
||||
row.get::<Option<String>, _>("notes"),
|
||||
row.get::<Option<String>, _>("photo_url"),
|
||||
row.get("birthday"),
|
||||
row.get("anniversary"),
|
||||
row.get("vcard"),
|
||||
row.get("etag"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
));
|
||||
}
|
||||
Ok(contacts)
|
||||
}
|
||||
|
||||
async fn get_groups_for_contact(&self, contact_id: &Uuid) -> ContactRepositoryResult<Vec<ContactGroup>> {
|
||||
let rows = sqlx::query(
|
||||
"SELECT g.id, g.address_book_id, g.name, g.created_at, g.updated_at FROM carddav.contact_groups g INNER JOIN carddav.group_memberships gm ON g.id = gm.group_id WHERE gm.contact_id = $1 ORDER BY g.name"
|
||||
)
|
||||
.bind(contact_id)
|
||||
.fetch_all(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to get groups for contact: {}", e)))?;
|
||||
|
||||
Ok(rows.into_iter().map(|row| {
|
||||
ContactGroup::from_raw(
|
||||
row.get::<Uuid, _>("id"),
|
||||
row.get::<Uuid, _>("address_book_id"),
|
||||
row.get::<String, _>("name"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
)
|
||||
}).collect())
|
||||
}
|
||||
}
|
||||
use async_trait::async_trait;
|
||||
use chrono::Utc;
|
||||
use serde_json::Value as JsonValue;
|
||||
use sqlx::{PgPool, Row, types::Uuid};
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::contact_persistence_dto::{
|
||||
AddressPersistenceDto, EmailPersistenceDto, PhonePersistenceDto, addresses_from_persistence,
|
||||
emails_from_persistence, phones_from_persistence,
|
||||
};
|
||||
use crate::common::errors::{DomainError, ErrorKind};
|
||||
use crate::domain::entities::contact::{Contact, ContactGroup};
|
||||
use crate::domain::repositories::contact_repository::{
|
||||
ContactGroupRepository, ContactRepositoryResult,
|
||||
};
|
||||
|
||||
pub struct ContactGroupPgRepository {
|
||||
pool: Arc<PgPool>,
|
||||
}
|
||||
|
||||
impl ContactGroupPgRepository {
|
||||
pub fn new(pool: Arc<PgPool>) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ContactGroupRepository for ContactGroupPgRepository {
|
||||
async fn create_group(&self, group: ContactGroup) -> ContactRepositoryResult<ContactGroup> {
|
||||
sqlx::query(
|
||||
"INSERT INTO carddav.contact_groups (id, address_book_id, name, created_at, updated_at) VALUES ($1, $2, $3, $4, $5)"
|
||||
)
|
||||
.bind(group.id())
|
||||
.bind(group.address_book_id())
|
||||
.bind(group.name())
|
||||
.bind(group.created_at())
|
||||
.bind(group.updated_at())
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to create group: {}", e)))?;
|
||||
|
||||
Ok(group)
|
||||
}
|
||||
|
||||
async fn update_group(&self, group: ContactGroup) -> ContactRepositoryResult<ContactGroup> {
|
||||
sqlx::query("UPDATE carddav.contact_groups SET name = $1, updated_at = $2 WHERE id = $3")
|
||||
.bind(group.name())
|
||||
.bind(Utc::now())
|
||||
.bind(group.id())
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"ContactGroup",
|
||||
format!("Failed to update group: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(group)
|
||||
}
|
||||
|
||||
async fn delete_group(&self, id: &Uuid) -> ContactRepositoryResult<()> {
|
||||
// Delete memberships first
|
||||
sqlx::query("DELETE FROM carddav.group_memberships WHERE group_id = $1")
|
||||
.bind(id)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"ContactGroup",
|
||||
format!("Failed to delete group memberships: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
sqlx::query("DELETE FROM carddav.contact_groups WHERE id = $1")
|
||||
.bind(id)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"ContactGroup",
|
||||
format!("Failed to delete group: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get_group_by_id(&self, id: &Uuid) -> ContactRepositoryResult<Option<ContactGroup>> {
|
||||
let row = sqlx::query(
|
||||
"SELECT id, address_book_id, name, created_at, updated_at FROM carddav.contact_groups WHERE id = $1"
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_optional(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to get group: {}", e)))?;
|
||||
|
||||
match row {
|
||||
Some(row) => {
|
||||
let group = ContactGroup::from_raw(
|
||||
row.get::<Uuid, _>("id"),
|
||||
row.get::<Uuid, _>("address_book_id"),
|
||||
row.get::<String, _>("name"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
);
|
||||
Ok(Some(group))
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_groups_by_address_book(
|
||||
&self,
|
||||
address_book_id: &Uuid,
|
||||
) -> ContactRepositoryResult<Vec<ContactGroup>> {
|
||||
let rows = sqlx::query(
|
||||
"SELECT id, address_book_id, name, created_at, updated_at FROM carddav.contact_groups WHERE address_book_id = $1 ORDER BY name"
|
||||
)
|
||||
.bind(address_book_id)
|
||||
.fetch_all(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to list groups: {}", e)))?;
|
||||
|
||||
Ok(rows
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
ContactGroup::from_raw(
|
||||
row.get::<Uuid, _>("id"),
|
||||
row.get::<Uuid, _>("address_book_id"),
|
||||
row.get::<String, _>("name"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
)
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn add_contact_to_group(
|
||||
&self,
|
||||
group_id: &Uuid,
|
||||
contact_id: &Uuid,
|
||||
) -> ContactRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
"INSERT INTO carddav.group_memberships (group_id, contact_id) VALUES ($1, $2) ON CONFLICT DO NOTHING"
|
||||
)
|
||||
.bind(group_id)
|
||||
.bind(contact_id)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to add contact to group: {}", e)))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_contact_from_group(
|
||||
&self,
|
||||
group_id: &Uuid,
|
||||
contact_id: &Uuid,
|
||||
) -> ContactRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
"DELETE FROM carddav.group_memberships WHERE group_id = $1 AND contact_id = $2",
|
||||
)
|
||||
.bind(group_id)
|
||||
.bind(contact_id)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"ContactGroup",
|
||||
format!("Failed to remove contact from group: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get_contacts_in_group(
|
||||
&self,
|
||||
group_id: &Uuid,
|
||||
) -> ContactRepositoryResult<Vec<Contact>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
c.id, c.address_book_id, c.uid, c.full_name, c.first_name, c.last_name, c.nickname,
|
||||
c.email, c.phone, c.address, c.organization, c.title, c.notes, c.photo_url,
|
||||
c.birthday, c.anniversary, c.vcard, c.etag, c.created_at, c.updated_at
|
||||
FROM carddav.contacts c
|
||||
INNER JOIN carddav.group_memberships gm ON c.id = gm.contact_id
|
||||
WHERE gm.group_id = $1
|
||||
ORDER BY c.full_name, c.first_name, c.last_name
|
||||
"#,
|
||||
)
|
||||
.bind(group_id)
|
||||
.fetch_all(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"ContactGroup",
|
||||
format!("Failed to get contacts in group: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
let mut contacts = Vec::new();
|
||||
for row in &rows {
|
||||
let email_json: JsonValue = row.get("email");
|
||||
let phone_json: JsonValue = row.get("phone");
|
||||
let address_json: JsonValue = row.get("address");
|
||||
|
||||
let emails = serde_json::from_value::<Vec<EmailPersistenceDto>>(email_json)
|
||||
.map(emails_from_persistence)
|
||||
.unwrap_or_default();
|
||||
let phones = serde_json::from_value::<Vec<PhonePersistenceDto>>(phone_json)
|
||||
.map(phones_from_persistence)
|
||||
.unwrap_or_default();
|
||||
let addresses = serde_json::from_value::<Vec<AddressPersistenceDto>>(address_json)
|
||||
.map(addresses_from_persistence)
|
||||
.unwrap_or_default();
|
||||
|
||||
contacts.push(Contact::from_raw(
|
||||
row.get("id"),
|
||||
row.get("address_book_id"),
|
||||
row.get("uid"),
|
||||
row.get::<Option<String>, _>("full_name"),
|
||||
row.get::<Option<String>, _>("first_name"),
|
||||
row.get::<Option<String>, _>("last_name"),
|
||||
row.get::<Option<String>, _>("nickname"),
|
||||
emails,
|
||||
phones,
|
||||
addresses,
|
||||
row.get::<Option<String>, _>("organization"),
|
||||
row.get::<Option<String>, _>("title"),
|
||||
row.get::<Option<String>, _>("notes"),
|
||||
row.get::<Option<String>, _>("photo_url"),
|
||||
row.get("birthday"),
|
||||
row.get("anniversary"),
|
||||
row.get("vcard"),
|
||||
row.get("etag"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
));
|
||||
}
|
||||
Ok(contacts)
|
||||
}
|
||||
|
||||
async fn get_groups_for_contact(
|
||||
&self,
|
||||
contact_id: &Uuid,
|
||||
) -> ContactRepositoryResult<Vec<ContactGroup>> {
|
||||
let rows = sqlx::query(
|
||||
"SELECT g.id, g.address_book_id, g.name, g.created_at, g.updated_at FROM carddav.contact_groups g INNER JOIN carddav.group_memberships gm ON g.id = gm.group_id WHERE gm.contact_id = $1 ORDER BY g.name"
|
||||
)
|
||||
.bind(contact_id)
|
||||
.fetch_all(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(ErrorKind::InternalError, "ContactGroup", format!("Failed to get groups for contact: {}", e)))?;
|
||||
|
||||
Ok(rows
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
ContactGroup::from_raw(
|
||||
row.get::<Uuid, _>("id"),
|
||||
row.get::<Uuid, _>("address_book_id"),
|
||||
row.get::<String, _>("name"),
|
||||
row.get("created_at"),
|
||||
row.get("updated_at"),
|
||||
)
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,129 +1,129 @@
|
||||
//! Persistence DTOs for Contact entities
|
||||
//!
|
||||
//! These DTOs are used for JSONB serialization/deserialization in PostgreSQL.
|
||||
//! They mirror the domain entities but include serde traits required for persistence.
|
||||
//! This keeps the domain layer free of infrastructure concerns (serde dependency).
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use crate::domain::entities::contact::{Email, Phone, Address};
|
||||
|
||||
/// Persistence DTO for Email - used for JSONB serialization
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct EmailPersistenceDto {
|
||||
pub email: String,
|
||||
pub r#type: String,
|
||||
pub is_primary: bool,
|
||||
}
|
||||
|
||||
impl From<&Email> for EmailPersistenceDto {
|
||||
fn from(email: &Email) -> Self {
|
||||
Self {
|
||||
email: email.email.clone(),
|
||||
r#type: email.r#type.clone(),
|
||||
is_primary: email.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<EmailPersistenceDto> for Email {
|
||||
fn from(dto: EmailPersistenceDto) -> Self {
|
||||
Self {
|
||||
email: dto.email,
|
||||
r#type: dto.r#type,
|
||||
is_primary: dto.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Persistence DTO for Phone - used for JSONB serialization
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct PhonePersistenceDto {
|
||||
pub number: String,
|
||||
pub r#type: String,
|
||||
pub is_primary: bool,
|
||||
}
|
||||
|
||||
impl From<&Phone> for PhonePersistenceDto {
|
||||
fn from(phone: &Phone) -> Self {
|
||||
Self {
|
||||
number: phone.number.clone(),
|
||||
r#type: phone.r#type.clone(),
|
||||
is_primary: phone.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<PhonePersistenceDto> for Phone {
|
||||
fn from(dto: PhonePersistenceDto) -> Self {
|
||||
Self {
|
||||
number: dto.number,
|
||||
r#type: dto.r#type,
|
||||
is_primary: dto.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Persistence DTO for Address - used for JSONB serialization
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct AddressPersistenceDto {
|
||||
pub street: Option<String>,
|
||||
pub city: Option<String>,
|
||||
pub state: Option<String>,
|
||||
pub postal_code: Option<String>,
|
||||
pub country: Option<String>,
|
||||
pub r#type: String,
|
||||
pub is_primary: bool,
|
||||
}
|
||||
|
||||
impl From<&Address> for AddressPersistenceDto {
|
||||
fn from(addr: &Address) -> Self {
|
||||
Self {
|
||||
street: addr.street.clone(),
|
||||
city: addr.city.clone(),
|
||||
state: addr.state.clone(),
|
||||
postal_code: addr.postal_code.clone(),
|
||||
country: addr.country.clone(),
|
||||
r#type: addr.r#type.clone(),
|
||||
is_primary: addr.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<AddressPersistenceDto> for Address {
|
||||
fn from(dto: AddressPersistenceDto) -> Self {
|
||||
Self {
|
||||
street: dto.street,
|
||||
city: dto.city,
|
||||
state: dto.state,
|
||||
postal_code: dto.postal_code,
|
||||
country: dto.country,
|
||||
r#type: dto.r#type,
|
||||
is_primary: dto.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper functions to convert collections
|
||||
pub fn emails_to_persistence(emails: &[Email]) -> Vec<EmailPersistenceDto> {
|
||||
emails.iter().map(EmailPersistenceDto::from).collect()
|
||||
}
|
||||
|
||||
pub fn emails_from_persistence(dtos: Vec<EmailPersistenceDto>) -> Vec<Email> {
|
||||
dtos.into_iter().map(Email::from).collect()
|
||||
}
|
||||
|
||||
pub fn phones_to_persistence(phones: &[Phone]) -> Vec<PhonePersistenceDto> {
|
||||
phones.iter().map(PhonePersistenceDto::from).collect()
|
||||
}
|
||||
|
||||
pub fn phones_from_persistence(dtos: Vec<PhonePersistenceDto>) -> Vec<Phone> {
|
||||
dtos.into_iter().map(Phone::from).collect()
|
||||
}
|
||||
|
||||
pub fn addresses_to_persistence(addresses: &[Address]) -> Vec<AddressPersistenceDto> {
|
||||
addresses.iter().map(AddressPersistenceDto::from).collect()
|
||||
}
|
||||
|
||||
pub fn addresses_from_persistence(dtos: Vec<AddressPersistenceDto>) -> Vec<Address> {
|
||||
dtos.into_iter().map(Address::from).collect()
|
||||
}
|
||||
//! Persistence DTOs for Contact entities
|
||||
//!
|
||||
//! These DTOs are used for JSONB serialization/deserialization in PostgreSQL.
|
||||
//! They mirror the domain entities but include serde traits required for persistence.
|
||||
//! This keeps the domain layer free of infrastructure concerns (serde dependency).
|
||||
|
||||
use crate::domain::entities::contact::{Address, Email, Phone};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// Persistence DTO for Email - used for JSONB serialization
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct EmailPersistenceDto {
|
||||
pub email: String,
|
||||
pub r#type: String,
|
||||
pub is_primary: bool,
|
||||
}
|
||||
|
||||
impl From<&Email> for EmailPersistenceDto {
|
||||
fn from(email: &Email) -> Self {
|
||||
Self {
|
||||
email: email.email.clone(),
|
||||
r#type: email.r#type.clone(),
|
||||
is_primary: email.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<EmailPersistenceDto> for Email {
|
||||
fn from(dto: EmailPersistenceDto) -> Self {
|
||||
Self {
|
||||
email: dto.email,
|
||||
r#type: dto.r#type,
|
||||
is_primary: dto.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Persistence DTO for Phone - used for JSONB serialization
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct PhonePersistenceDto {
|
||||
pub number: String,
|
||||
pub r#type: String,
|
||||
pub is_primary: bool,
|
||||
}
|
||||
|
||||
impl From<&Phone> for PhonePersistenceDto {
|
||||
fn from(phone: &Phone) -> Self {
|
||||
Self {
|
||||
number: phone.number.clone(),
|
||||
r#type: phone.r#type.clone(),
|
||||
is_primary: phone.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<PhonePersistenceDto> for Phone {
|
||||
fn from(dto: PhonePersistenceDto) -> Self {
|
||||
Self {
|
||||
number: dto.number,
|
||||
r#type: dto.r#type,
|
||||
is_primary: dto.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Persistence DTO for Address - used for JSONB serialization
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct AddressPersistenceDto {
|
||||
pub street: Option<String>,
|
||||
pub city: Option<String>,
|
||||
pub state: Option<String>,
|
||||
pub postal_code: Option<String>,
|
||||
pub country: Option<String>,
|
||||
pub r#type: String,
|
||||
pub is_primary: bool,
|
||||
}
|
||||
|
||||
impl From<&Address> for AddressPersistenceDto {
|
||||
fn from(addr: &Address) -> Self {
|
||||
Self {
|
||||
street: addr.street.clone(),
|
||||
city: addr.city.clone(),
|
||||
state: addr.state.clone(),
|
||||
postal_code: addr.postal_code.clone(),
|
||||
country: addr.country.clone(),
|
||||
r#type: addr.r#type.clone(),
|
||||
is_primary: addr.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<AddressPersistenceDto> for Address {
|
||||
fn from(dto: AddressPersistenceDto) -> Self {
|
||||
Self {
|
||||
street: dto.street,
|
||||
city: dto.city,
|
||||
state: dto.state,
|
||||
postal_code: dto.postal_code,
|
||||
country: dto.country,
|
||||
r#type: dto.r#type,
|
||||
is_primary: dto.is_primary,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper functions to convert collections
|
||||
pub fn emails_to_persistence(emails: &[Email]) -> Vec<EmailPersistenceDto> {
|
||||
emails.iter().map(EmailPersistenceDto::from).collect()
|
||||
}
|
||||
|
||||
pub fn emails_from_persistence(dtos: Vec<EmailPersistenceDto>) -> Vec<Email> {
|
||||
dtos.into_iter().map(Email::from).collect()
|
||||
}
|
||||
|
||||
pub fn phones_to_persistence(phones: &[Phone]) -> Vec<PhonePersistenceDto> {
|
||||
phones.iter().map(PhonePersistenceDto::from).collect()
|
||||
}
|
||||
|
||||
pub fn phones_from_persistence(dtos: Vec<PhonePersistenceDto>) -> Vec<Phone> {
|
||||
dtos.into_iter().map(Phone::from).collect()
|
||||
}
|
||||
|
||||
pub fn addresses_to_persistence(addresses: &[Address]) -> Vec<AddressPersistenceDto> {
|
||||
addresses.iter().map(AddressPersistenceDto::from).collect()
|
||||
}
|
||||
|
||||
pub fn addresses_from_persistence(dtos: Vec<AddressPersistenceDto>) -> Vec<Address> {
|
||||
dtos.into_iter().map(Address::from).collect()
|
||||
}
|
||||
|
||||
@@ -1,17 +1,17 @@
|
||||
use async_trait::async_trait;
|
||||
use chrono::Utc;
|
||||
use serde_json::Value as JsonValue;
|
||||
use sqlx::{PgPool, Row, types::Uuid};
|
||||
use std::sync::Arc;
|
||||
use serde_json::Value as JsonValue;
|
||||
|
||||
use super::contact_persistence_dto::{
|
||||
AddressPersistenceDto, EmailPersistenceDto, PhonePersistenceDto, addresses_from_persistence,
|
||||
addresses_to_persistence, emails_from_persistence, emails_to_persistence,
|
||||
phones_from_persistence, phones_to_persistence,
|
||||
};
|
||||
use crate::common::errors::DomainError;
|
||||
use crate::domain::entities::contact::Contact;
|
||||
use crate::domain::repositories::contact_repository::{ContactRepository, ContactRepositoryResult};
|
||||
use crate::common::errors::DomainError;
|
||||
use super::contact_persistence_dto::{
|
||||
emails_to_persistence, phones_to_persistence, addresses_to_persistence,
|
||||
emails_from_persistence, phones_from_persistence, addresses_from_persistence,
|
||||
EmailPersistenceDto, PhonePersistenceDto, AddressPersistenceDto,
|
||||
};
|
||||
|
||||
pub struct ContactPgRepository {
|
||||
pool: Arc<PgPool>,
|
||||
@@ -70,11 +70,11 @@ impl ContactRepository for ContactPgRepository {
|
||||
let email_dtos = emails_to_persistence(contact.email());
|
||||
let phone_dtos = phones_to_persistence(contact.phone());
|
||||
let address_dtos = addresses_to_persistence(contact.address());
|
||||
|
||||
|
||||
let email_json = serde_json::to_value(&email_dtos).unwrap_or(JsonValue::Null);
|
||||
let phone_json = serde_json::to_value(&phone_dtos).unwrap_or(JsonValue::Null);
|
||||
let address_json = serde_json::to_value(&address_dtos).unwrap_or(JsonValue::Null);
|
||||
|
||||
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
INSERT INTO carddav.contacts (
|
||||
@@ -90,7 +90,7 @@ impl ContactRepository for ContactPgRepository {
|
||||
id, address_book_id, uid, full_name, first_name, last_name, nickname,
|
||||
email, phone, address, organization, title, notes, photo_url,
|
||||
birthday, anniversary, vcard, etag, created_at, updated_at
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(contact.id())
|
||||
.bind(contact.address_book_id())
|
||||
@@ -125,15 +125,15 @@ impl ContactRepository for ContactPgRepository {
|
||||
let email_dtos = emails_to_persistence(contact.email());
|
||||
let phone_dtos = phones_to_persistence(contact.phone());
|
||||
let address_dtos = addresses_to_persistence(contact.address());
|
||||
|
||||
|
||||
let email_json = serde_json::to_value(&email_dtos).unwrap_or(JsonValue::Null);
|
||||
let phone_json = serde_json::to_value(&phone_dtos).unwrap_or(JsonValue::Null);
|
||||
let address_json = serde_json::to_value(&address_dtos).unwrap_or(JsonValue::Null);
|
||||
|
||||
|
||||
// Create a clone of the contact with the updated timestamp
|
||||
let mut updated_contact = contact.clone();
|
||||
updated_contact.set_updated_at(now);
|
||||
|
||||
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
UPDATE carddav.contacts
|
||||
@@ -159,7 +159,7 @@ impl ContactRepository for ContactPgRepository {
|
||||
id, address_book_id, uid, full_name, first_name, last_name, nickname,
|
||||
email, phone, address, organization, title, notes, photo_url,
|
||||
birthday, anniversary, vcard, etag, created_at, updated_at
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(updated_contact.full_name_owned())
|
||||
.bind(updated_contact.first_name_owned())
|
||||
@@ -190,7 +190,7 @@ impl ContactRepository for ContactPgRepository {
|
||||
r#"
|
||||
DELETE FROM carddav.contacts
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(id)
|
||||
.execute(&*self.pool)
|
||||
@@ -209,7 +209,7 @@ impl ContactRepository for ContactPgRepository {
|
||||
birthday, anniversary, vcard, etag, created_at, updated_at
|
||||
FROM carddav.contacts
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_optional(&*self.pool)
|
||||
@@ -222,7 +222,11 @@ impl ContactRepository for ContactPgRepository {
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_contact_by_uid(&self, address_book_id: &Uuid, uid: &str) -> ContactRepositoryResult<Option<Contact>> {
|
||||
async fn get_contact_by_uid(
|
||||
&self,
|
||||
address_book_id: &Uuid,
|
||||
uid: &str,
|
||||
) -> ContactRepositoryResult<Option<Contact>> {
|
||||
let row_opt = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -231,7 +235,7 @@ impl ContactRepository for ContactPgRepository {
|
||||
birthday, anniversary, vcard, etag, created_at, updated_at
|
||||
FROM carddav.contacts
|
||||
WHERE address_book_id = $1 AND uid = $2
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(address_book_id)
|
||||
.bind(uid)
|
||||
@@ -245,7 +249,10 @@ impl ContactRepository for ContactPgRepository {
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_contacts_by_address_book(&self, address_book_id: &Uuid) -> ContactRepositoryResult<Vec<Contact>> {
|
||||
async fn get_contacts_by_address_book(
|
||||
&self,
|
||||
address_book_id: &Uuid,
|
||||
) -> ContactRepositoryResult<Vec<Contact>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -255,12 +262,14 @@ impl ContactRepository for ContactPgRepository {
|
||||
FROM carddav.contacts
|
||||
WHERE address_book_id = $1
|
||||
ORDER BY full_name, first_name, last_name
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(address_book_id)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get contacts by address book: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get contacts by address book: {}", e))
|
||||
})?;
|
||||
|
||||
let mut contacts = Vec::new();
|
||||
for row in &rows {
|
||||
@@ -271,7 +280,7 @@ impl ContactRepository for ContactPgRepository {
|
||||
|
||||
async fn get_contacts_by_email(&self, email: &str) -> ContactRepositoryResult<Vec<Contact>> {
|
||||
let search_pattern = format!("%{}%", email);
|
||||
|
||||
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -281,12 +290,14 @@ impl ContactRepository for ContactPgRepository {
|
||||
FROM carddav.contacts
|
||||
WHERE email::text ILIKE $1
|
||||
ORDER BY full_name, first_name, last_name
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(&search_pattern)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get contacts by email: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get contacts by email: {}", e))
|
||||
})?;
|
||||
|
||||
let mut contacts = Vec::new();
|
||||
for row in &rows {
|
||||
@@ -295,7 +306,10 @@ impl ContactRepository for ContactPgRepository {
|
||||
Ok(contacts)
|
||||
}
|
||||
|
||||
async fn get_contacts_by_group(&self, group_id: &Uuid) -> ContactRepositoryResult<Vec<Contact>> {
|
||||
async fn get_contacts_by_group(
|
||||
&self,
|
||||
group_id: &Uuid,
|
||||
) -> ContactRepositoryResult<Vec<Contact>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -306,12 +320,14 @@ impl ContactRepository for ContactPgRepository {
|
||||
INNER JOIN carddav.group_memberships m ON c.id = m.contact_id
|
||||
WHERE m.group_id = $1
|
||||
ORDER BY c.full_name, c.first_name, c.last_name
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(group_id)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(|e| DomainError::database_error(format!("Failed to get contacts by group: {}", e)))?;
|
||||
.map_err(|e| {
|
||||
DomainError::database_error(format!("Failed to get contacts by group: {}", e))
|
||||
})?;
|
||||
|
||||
let mut contacts = Vec::new();
|
||||
for row in &rows {
|
||||
@@ -320,9 +336,13 @@ impl ContactRepository for ContactPgRepository {
|
||||
Ok(contacts)
|
||||
}
|
||||
|
||||
async fn search_contacts(&self, address_book_id: &Uuid, query: &str) -> ContactRepositoryResult<Vec<Contact>> {
|
||||
async fn search_contacts(
|
||||
&self,
|
||||
address_book_id: &Uuid,
|
||||
query: &str,
|
||||
) -> ContactRepositoryResult<Vec<Contact>> {
|
||||
let search_pattern = format!("%{}%", query);
|
||||
|
||||
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -341,7 +361,7 @@ impl ContactRepository for ContactPgRepository {
|
||||
OR organization ILIKE $2
|
||||
)
|
||||
ORDER BY full_name, first_name, last_name
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(address_book_id)
|
||||
.bind(&search_pattern)
|
||||
@@ -355,4 +375,4 @@ impl ContactRepository for ContactPgRepository {
|
||||
}
|
||||
Ok(contacts)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
use std::sync::Arc;
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{PgPool, Row};
|
||||
use std::sync::Arc;
|
||||
use tracing::error;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::application::dtos::favorites_dto::FavoriteItemDto;
|
||||
use crate::application::ports::favorites_ports::FavoritesRepositoryPort;
|
||||
use crate::common::errors::{Result, DomainError, ErrorKind};
|
||||
use crate::common::errors::{DomainError, ErrorKind, Result};
|
||||
|
||||
/// PostgreSQL implementation of the favorites persistence port.
|
||||
pub struct FavoritesPgRepository {
|
||||
@@ -42,7 +42,11 @@ impl FavoritesRepositoryPort for FavoritesPgRepository {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!("Database error fetching favorites: {}", e);
|
||||
DomainError::new(ErrorKind::InternalError, "Favorites", format!("Failed to fetch favorites: {}", e))
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Favorites",
|
||||
format!("Failed to fetch favorites: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
let favorites = rows
|
||||
@@ -76,7 +80,11 @@ impl FavoritesRepositoryPort for FavoritesPgRepository {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!("Database error adding favorite: {}", e);
|
||||
DomainError::new(ErrorKind::InternalError, "Favorites", format!("Failed to add to favorites: {}", e))
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Favorites",
|
||||
format!("Failed to add to favorites: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
@@ -98,7 +106,11 @@ impl FavoritesRepositoryPort for FavoritesPgRepository {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!("Database error removing favorite: {}", e);
|
||||
DomainError::new(ErrorKind::InternalError, "Favorites", format!("Failed to remove from favorites: {}", e))
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Favorites",
|
||||
format!("Failed to remove from favorites: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(result.rows_affected() > 0)
|
||||
@@ -122,7 +134,11 @@ impl FavoritesRepositoryPort for FavoritesPgRepository {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!("Database error checking favorite status: {}", e);
|
||||
DomainError::new(ErrorKind::InternalError, "Favorites", format!("Failed to check favorite status: {}", e))
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Favorites",
|
||||
format!("Failed to check favorite status: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(row.try_get("is_favorite").unwrap_or(false))
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
mod address_book_pg_repository;
|
||||
mod calendar_pg_repository;
|
||||
mod calendar_event_pg_repository;
|
||||
mod contact_pg_repository;
|
||||
mod calendar_pg_repository;
|
||||
mod contact_group_pg_repository;
|
||||
mod contact_persistence_dto;
|
||||
mod contact_pg_repository;
|
||||
mod favorites_pg_repository;
|
||||
mod recent_items_pg_repository;
|
||||
mod session_pg_repository;
|
||||
@@ -12,11 +12,11 @@ mod transaction_utils;
|
||||
mod user_pg_repository;
|
||||
|
||||
pub use address_book_pg_repository::AddressBookPgRepository;
|
||||
pub use calendar_pg_repository::CalendarPgRepository;
|
||||
pub use calendar_event_pg_repository::CalendarEventPgRepository;
|
||||
pub use contact_pg_repository::ContactPgRepository;
|
||||
pub use calendar_pg_repository::CalendarPgRepository;
|
||||
pub use contact_group_pg_repository::ContactGroupPgRepository;
|
||||
pub use contact_persistence_dto::*;
|
||||
pub use contact_pg_repository::ContactPgRepository;
|
||||
pub use favorites_pg_repository::FavoritesPgRepository;
|
||||
pub use recent_items_pg_repository::RecentItemsPgRepository;
|
||||
pub use session_pg_repository::SessionPgRepository;
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
use std::sync::Arc;
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{PgPool, Row};
|
||||
use std::sync::Arc;
|
||||
use tracing::error;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::application::dtos::recent_dto::RecentItemDto;
|
||||
use crate::application::ports::recent_ports::RecentItemsRepositoryPort;
|
||||
use crate::common::errors::{Result, DomainError, ErrorKind};
|
||||
use crate::common::errors::{DomainError, ErrorKind, Result};
|
||||
|
||||
/// PostgreSQL implementation of the recent items persistence port.
|
||||
pub struct RecentItemsPgRepository {
|
||||
@@ -44,7 +44,11 @@ impl RecentItemsRepositoryPort for RecentItemsPgRepository {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!("Database error fetching recent items: {}", e);
|
||||
DomainError::new(ErrorKind::InternalError, "RecentItems", format!("Failed to fetch recent items: {}", e))
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"RecentItems",
|
||||
format!("Failed to fetch recent items: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
let items = rows
|
||||
@@ -79,7 +83,11 @@ impl RecentItemsRepositoryPort for RecentItemsPgRepository {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!("Database error upserting recent item access: {}", e);
|
||||
DomainError::new(ErrorKind::InternalError, "RecentItems", format!("Failed to record item access: {}", e))
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"RecentItems",
|
||||
format!("Failed to record item access: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
@@ -101,7 +109,11 @@ impl RecentItemsRepositoryPort for RecentItemsPgRepository {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!("Database error removing recent item: {}", e);
|
||||
DomainError::new(ErrorKind::InternalError, "RecentItems", format!("Failed to remove recent item: {}", e))
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"RecentItems",
|
||||
format!("Failed to remove recent item: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(result.rows_affected() > 0)
|
||||
@@ -121,7 +133,11 @@ impl RecentItemsRepositoryPort for RecentItemsPgRepository {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!("Database error clearing recent items: {}", e);
|
||||
DomainError::new(ErrorKind::InternalError, "RecentItems", format!("Failed to clear recent items: {}", e))
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"RecentItems",
|
||||
format!("Failed to clear recent items: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
@@ -147,7 +163,11 @@ impl RecentItemsRepositoryPort for RecentItemsPgRepository {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!("Database error pruning old recent items: {}", e);
|
||||
DomainError::new(ErrorKind::InternalError, "RecentItems", format!("Failed to prune recent items: {}", e))
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"RecentItems",
|
||||
format!("Failed to prune recent items: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
|
||||
@@ -1,13 +1,15 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{PgPool, Row};
|
||||
use std::sync::Arc;
|
||||
use chrono::Utc;
|
||||
use futures::future::BoxFuture;
|
||||
use sqlx::{PgPool, Row};
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::domain::entities::session::Session;
|
||||
use crate::domain::repositories::session_repository::{SessionRepository, SessionRepositoryError, SessionRepositoryResult};
|
||||
use crate::application::ports::auth_ports::SessionStoragePort;
|
||||
use crate::common::errors::DomainError;
|
||||
use crate::domain::entities::session::Session;
|
||||
use crate::domain::repositories::session_repository::{
|
||||
SessionRepository, SessionRepositoryError, SessionRepositoryResult,
|
||||
};
|
||||
use crate::infrastructure::repositories::pg::transaction_utils::with_transaction;
|
||||
|
||||
// Implement From<sqlx::Error> for SessionRepositoryError to allow automatic conversions
|
||||
@@ -25,16 +27,14 @@ impl SessionPgRepository {
|
||||
pub fn new(pool: Arc<PgPool>) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
|
||||
// Helper method to map SQL errors to domain errors
|
||||
pub fn map_sqlx_error(err: sqlx::Error) -> SessionRepositoryError {
|
||||
match err {
|
||||
sqlx::Error::RowNotFound => {
|
||||
SessionRepositoryError::NotFound("Session not found".to_string())
|
||||
},
|
||||
_ => SessionRepositoryError::DatabaseError(
|
||||
format!("Database error: {}", err)
|
||||
),
|
||||
}
|
||||
_ => SessionRepositoryError::DatabaseError(format!("Database error: {}", err)),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -45,65 +45,66 @@ impl SessionRepository for SessionPgRepository {
|
||||
async fn create_session(&self, session: Session) -> SessionRepositoryResult<Session> {
|
||||
// Create a copy of the session for the closure
|
||||
let session_clone = session.clone();
|
||||
|
||||
with_transaction(
|
||||
&self.pool,
|
||||
"create_session",
|
||||
|tx| {
|
||||
Box::pin(async move {
|
||||
// Insert the session
|
||||
sqlx::query(
|
||||
r#"
|
||||
|
||||
with_transaction(&self.pool, "create_session", |tx| {
|
||||
Box::pin(async move {
|
||||
// Insert the session
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO auth.sessions (
|
||||
id, user_id, refresh_token, expires_at,
|
||||
ip_address, user_agent, created_at, revoked
|
||||
) VALUES (
|
||||
$1, $2, $3, $4, $5, $6, $7, $8
|
||||
)
|
||||
"#
|
||||
)
|
||||
.bind(session_clone.id())
|
||||
.bind(session_clone.user_id())
|
||||
.bind(session_clone.refresh_token())
|
||||
.bind(session_clone.expires_at())
|
||||
.bind(session_clone.ip_address())
|
||||
.bind(session_clone.user_agent())
|
||||
.bind(session_clone.created_at())
|
||||
.bind(session_clone.is_revoked())
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
// Optionally, update the user's last login
|
||||
// within the same transaction
|
||||
sqlx::query(
|
||||
r#"
|
||||
"#,
|
||||
)
|
||||
.bind(session_clone.id())
|
||||
.bind(session_clone.user_id())
|
||||
.bind(session_clone.refresh_token())
|
||||
.bind(session_clone.expires_at())
|
||||
.bind(session_clone.ip_address())
|
||||
.bind(session_clone.user_agent())
|
||||
.bind(session_clone.created_at())
|
||||
.bind(session_clone.is_revoked())
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
// Optionally, update the user's last login
|
||||
// within the same transaction
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE auth.users
|
||||
SET last_login_at = NOW(), updated_at = NOW()
|
||||
WHERE id = $1
|
||||
"#
|
||||
)
|
||||
.bind(session_clone.user_id())
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
// Convert the error but without interrupting session
|
||||
// creation if the update fails
|
||||
tracing::warn!("Could not update last_login_at for user {}: {}",
|
||||
session_clone.user_id(), e);
|
||||
SessionRepositoryError::DatabaseError(format!(
|
||||
"Session created but could not update last_login_at: {}", e
|
||||
))
|
||||
})?;
|
||||
|
||||
Ok(session_clone)
|
||||
}) as BoxFuture<'_, SessionRepositoryResult<Session>>
|
||||
}
|
||||
).await?;
|
||||
|
||||
"#,
|
||||
)
|
||||
.bind(session_clone.user_id())
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
// Convert the error but without interrupting session
|
||||
// creation if the update fails
|
||||
tracing::warn!(
|
||||
"Could not update last_login_at for user {}: {}",
|
||||
session_clone.user_id(),
|
||||
e
|
||||
);
|
||||
SessionRepositoryError::DatabaseError(format!(
|
||||
"Session created but could not update last_login_at: {}",
|
||||
e
|
||||
))
|
||||
})?;
|
||||
|
||||
Ok(session_clone)
|
||||
}) as BoxFuture<'_, SessionRepositoryResult<Session>>
|
||||
})
|
||||
.await?;
|
||||
|
||||
Ok(session)
|
||||
}
|
||||
|
||||
|
||||
/// Gets a session by ID
|
||||
async fn get_session_by_id(&self, id: &str) -> SessionRepositoryResult<Session> {
|
||||
let row = sqlx::query(
|
||||
@@ -113,7 +114,7 @@ impl SessionRepository for SessionPgRepository {
|
||||
ip_address, user_agent, created_at, revoked
|
||||
FROM auth.sessions
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_one(&*self.pool)
|
||||
@@ -131,9 +132,12 @@ impl SessionRepository for SessionPgRepository {
|
||||
row.get("revoked"),
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
/// Gets a session by refresh token
|
||||
async fn get_session_by_refresh_token(&self, refresh_token: &str) -> SessionRepositoryResult<Session> {
|
||||
async fn get_session_by_refresh_token(
|
||||
&self,
|
||||
refresh_token: &str,
|
||||
) -> SessionRepositoryResult<Session> {
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -141,7 +145,7 @@ impl SessionRepository for SessionPgRepository {
|
||||
ip_address, user_agent, created_at, revoked
|
||||
FROM auth.sessions
|
||||
WHERE refresh_token = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(refresh_token)
|
||||
.fetch_one(&*self.pool)
|
||||
@@ -159,9 +163,12 @@ impl SessionRepository for SessionPgRepository {
|
||||
row.get("revoked"),
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
/// Gets all sessions for a user
|
||||
async fn get_sessions_by_user_id(&self, user_id: &str) -> SessionRepositoryResult<Vec<Session>> {
|
||||
async fn get_sessions_by_user_id(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> SessionRepositoryResult<Vec<Session>> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -170,14 +177,15 @@ impl SessionRepository for SessionPgRepository {
|
||||
FROM auth.sessions
|
||||
WHERE user_id = $1
|
||||
ORDER BY created_at DESC
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
let sessions = rows.into_iter()
|
||||
let sessions = rows
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
Session::from_raw(
|
||||
row.get("id"),
|
||||
@@ -194,90 +202,84 @@ impl SessionRepository for SessionPgRepository {
|
||||
|
||||
Ok(sessions)
|
||||
}
|
||||
|
||||
|
||||
/// Revokes a specific session using a transaction
|
||||
async fn revoke_session(&self, session_id: &str) -> SessionRepositoryResult<()> {
|
||||
let id = session_id.to_string(); // Clone for use in closure
|
||||
|
||||
with_transaction(
|
||||
&self.pool,
|
||||
"revoke_session",
|
||||
|tx| {
|
||||
Box::pin(async move {
|
||||
// Revoke the session
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
|
||||
with_transaction(&self.pool, "revoke_session", |tx| {
|
||||
Box::pin(async move {
|
||||
// Revoke the session
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE auth.sessions
|
||||
SET revoked = true
|
||||
WHERE id = $1
|
||||
RETURNING user_id
|
||||
"#
|
||||
)
|
||||
.bind(&id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
// If we found the session, we can log a security event
|
||||
if let Some(row) = result {
|
||||
let user_id: String = row.try_get("user_id").unwrap_or_default();
|
||||
|
||||
// Log security event (in a security table)
|
||||
// This is optional but shows how additional operations
|
||||
// can be performed in the same transaction
|
||||
tracing::info!("Session with ID {} for user {} revoked", id, user_id);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}) as BoxFuture<'_, SessionRepositoryResult<()>>
|
||||
}
|
||||
).await
|
||||
"#,
|
||||
)
|
||||
.bind(&id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
// If we found the session, we can log a security event
|
||||
if let Some(row) = result {
|
||||
let user_id: String = row.try_get("user_id").unwrap_or_default();
|
||||
|
||||
// Log security event (in a security table)
|
||||
// This is optional but shows how additional operations
|
||||
// can be performed in the same transaction
|
||||
tracing::info!("Session with ID {} for user {} revoked", id, user_id);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}) as BoxFuture<'_, SessionRepositoryResult<()>>
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
|
||||
/// Revokes all sessions for a user using a transaction
|
||||
async fn revoke_all_user_sessions(&self, user_id: &str) -> SessionRepositoryResult<u64> {
|
||||
let user_id_clone = user_id.to_string(); // Clone for use in closure
|
||||
|
||||
with_transaction(
|
||||
&self.pool,
|
||||
"revoke_all_user_sessions",
|
||||
|tx| {
|
||||
Box::pin(async move {
|
||||
// Revoke all sessions for the user
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
|
||||
with_transaction(&self.pool, "revoke_all_user_sessions", |tx| {
|
||||
Box::pin(async move {
|
||||
// Revoke all sessions for the user
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE auth.sessions
|
||||
SET revoked = true
|
||||
WHERE user_id = $1 AND revoked = false
|
||||
"#
|
||||
)
|
||||
.bind(&user_id_clone)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
let affected = result.rows_affected();
|
||||
|
||||
// Log security event
|
||||
if affected > 0 {
|
||||
tracing::info!("Revoked {} sessions for user {}", affected, user_id_clone);
|
||||
}
|
||||
|
||||
Ok(affected)
|
||||
}) as BoxFuture<'_, SessionRepositoryResult<u64>>
|
||||
}
|
||||
).await
|
||||
"#,
|
||||
)
|
||||
.bind(&user_id_clone)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
let affected = result.rows_affected();
|
||||
|
||||
// Log security event
|
||||
if affected > 0 {
|
||||
tracing::info!("Revoked {} sessions for user {}", affected, user_id_clone);
|
||||
}
|
||||
|
||||
Ok(affected)
|
||||
}) as BoxFuture<'_, SessionRepositoryResult<u64>>
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
|
||||
/// Deletes expired sessions
|
||||
async fn delete_expired_sessions(&self) -> SessionRepositoryResult<u64> {
|
||||
let now = Utc::now();
|
||||
|
||||
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
DELETE FROM auth.sessions
|
||||
WHERE expires_at < $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(now)
|
||||
.execute(&*self.pool)
|
||||
@@ -292,22 +294,29 @@ impl SessionRepository for SessionPgRepository {
|
||||
#[async_trait]
|
||||
impl SessionStoragePort for SessionPgRepository {
|
||||
async fn create_session(&self, session: Session) -> Result<Session, DomainError> {
|
||||
SessionRepository::create_session(self, session).await.map_err(DomainError::from)
|
||||
SessionRepository::create_session(self, session)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
async fn get_session_by_refresh_token(&self, refresh_token: &str) -> Result<Session, DomainError> {
|
||||
|
||||
async fn get_session_by_refresh_token(
|
||||
&self,
|
||||
refresh_token: &str,
|
||||
) -> Result<Session, DomainError> {
|
||||
SessionRepository::get_session_by_refresh_token(self, refresh_token)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
|
||||
async fn revoke_session(&self, session_id: &str) -> Result<(), DomainError> {
|
||||
SessionRepository::revoke_session(self, session_id).await.map_err(DomainError::from)
|
||||
SessionRepository::revoke_session(self, session_id)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
|
||||
async fn revoke_all_user_sessions(&self, user_id: &str) -> Result<u64, DomainError> {
|
||||
SessionRepository::revoke_all_user_sessions(self, user_id)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,88 +1,102 @@
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use async_trait::async_trait;
|
||||
use sqlx::PgPool;
|
||||
|
||||
use crate::domain::repositories::settings_repository::SettingsRepository;
|
||||
use crate::common::errors::{DomainError, ErrorKind};
|
||||
|
||||
pub struct SettingsPgRepository {
|
||||
pool: Arc<PgPool>,
|
||||
}
|
||||
|
||||
impl SettingsPgRepository {
|
||||
pub fn new(pool: Arc<PgPool>) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SettingsRepository for SettingsPgRepository {
|
||||
async fn get(&self, key: &str) -> Result<Option<String>, DomainError> {
|
||||
let row = sqlx::query_scalar::<_, String>(
|
||||
"SELECT value FROM auth.admin_settings WHERE key = $1"
|
||||
)
|
||||
.bind(key)
|
||||
.fetch_optional(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(
|
||||
ErrorKind::InternalError, "Settings", format!("DB error: {}", e),
|
||||
))?;
|
||||
|
||||
Ok(row)
|
||||
}
|
||||
|
||||
async fn get_by_category(&self, category: &str) -> Result<HashMap<String, String>, DomainError> {
|
||||
let rows = sqlx::query_as::<_, (String, String)>(
|
||||
"SELECT key, value FROM auth.admin_settings WHERE category = $1"
|
||||
)
|
||||
.bind(category)
|
||||
.fetch_all(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(
|
||||
ErrorKind::InternalError, "Settings", format!("DB error: {}", e),
|
||||
))?;
|
||||
|
||||
Ok(rows.into_iter().collect())
|
||||
}
|
||||
|
||||
async fn set(
|
||||
&self,
|
||||
key: &str,
|
||||
value: &str,
|
||||
category: &str,
|
||||
is_secret: bool,
|
||||
updated_by: Option<&str>,
|
||||
) -> Result<(), DomainError> {
|
||||
sqlx::query(
|
||||
"INSERT INTO auth.admin_settings (key, value, category, is_secret, updated_by, updated_at)
|
||||
VALUES ($1, $2, $3, $4, $5, NOW())
|
||||
ON CONFLICT (key) DO UPDATE
|
||||
SET value = $2, category = $3, is_secret = $4, updated_by = $5, updated_at = NOW()"
|
||||
)
|
||||
.bind(key)
|
||||
.bind(value)
|
||||
.bind(category)
|
||||
.bind(is_secret)
|
||||
.bind(updated_by)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(
|
||||
ErrorKind::InternalError, "Settings", format!("DB error: {}", e),
|
||||
))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn delete(&self, key: &str) -> Result<(), DomainError> {
|
||||
sqlx::query("DELETE FROM auth.admin_settings WHERE key = $1")
|
||||
.bind(key)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(
|
||||
ErrorKind::InternalError, "Settings", format!("DB error: {}", e),
|
||||
))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
use async_trait::async_trait;
|
||||
use sqlx::PgPool;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::common::errors::{DomainError, ErrorKind};
|
||||
use crate::domain::repositories::settings_repository::SettingsRepository;
|
||||
|
||||
pub struct SettingsPgRepository {
|
||||
pool: Arc<PgPool>,
|
||||
}
|
||||
|
||||
impl SettingsPgRepository {
|
||||
pub fn new(pool: Arc<PgPool>) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SettingsRepository for SettingsPgRepository {
|
||||
async fn get(&self, key: &str) -> Result<Option<String>, DomainError> {
|
||||
let row =
|
||||
sqlx::query_scalar::<_, String>("SELECT value FROM auth.admin_settings WHERE key = $1")
|
||||
.bind(key)
|
||||
.fetch_optional(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Settings",
|
||||
format!("DB error: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(row)
|
||||
}
|
||||
|
||||
async fn get_by_category(
|
||||
&self,
|
||||
category: &str,
|
||||
) -> Result<HashMap<String, String>, DomainError> {
|
||||
let rows = sqlx::query_as::<_, (String, String)>(
|
||||
"SELECT key, value FROM auth.admin_settings WHERE category = $1",
|
||||
)
|
||||
.bind(category)
|
||||
.fetch_all(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Settings",
|
||||
format!("DB error: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(rows.into_iter().collect())
|
||||
}
|
||||
|
||||
async fn set(
|
||||
&self,
|
||||
key: &str,
|
||||
value: &str,
|
||||
category: &str,
|
||||
is_secret: bool,
|
||||
updated_by: Option<&str>,
|
||||
) -> Result<(), DomainError> {
|
||||
sqlx::query(
|
||||
"INSERT INTO auth.admin_settings (key, value, category, is_secret, updated_by, updated_at)
|
||||
VALUES ($1, $2, $3, $4, $5, NOW())
|
||||
ON CONFLICT (key) DO UPDATE
|
||||
SET value = $2, category = $3, is_secret = $4, updated_by = $5, updated_at = NOW()"
|
||||
)
|
||||
.bind(key)
|
||||
.bind(value)
|
||||
.bind(category)
|
||||
.bind(is_secret)
|
||||
.bind(updated_by)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| DomainError::new(
|
||||
ErrorKind::InternalError, "Settings", format!("DB error: {}", e),
|
||||
))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn delete(&self, key: &str) -> Result<(), DomainError> {
|
||||
sqlx::query("DELETE FROM auth.admin_settings WHERE key = $1")
|
||||
.bind(key)
|
||||
.execute(self.pool.as_ref())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Settings",
|
||||
format!("DB error: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use sqlx::{PgPool, Transaction, Postgres, Error as SqlxError};
|
||||
use sqlx::{Error as SqlxError, PgPool, Postgres, Transaction};
|
||||
use std::sync::Arc;
|
||||
use tracing::{debug, error, info};
|
||||
|
||||
@@ -13,17 +13,19 @@ pub async fn with_transaction<F, T, E>(
|
||||
operation: F,
|
||||
) -> Result<T, E>
|
||||
where
|
||||
F: for<'c> FnOnce(&'c mut Transaction<'_, Postgres>) -> futures::future::BoxFuture<'c, Result<T, E>>,
|
||||
F: for<'c> FnOnce(
|
||||
&'c mut Transaction<'_, Postgres>,
|
||||
) -> futures::future::BoxFuture<'c, Result<T, E>>,
|
||||
E: From<SqlxError> + std::fmt::Display,
|
||||
{
|
||||
debug!("Starting database transaction for: {}", operation_name);
|
||||
|
||||
|
||||
// Begin transaction
|
||||
let mut tx = pool.begin().await.map_err(|e| {
|
||||
error!("Failed to begin transaction for {}: {}", operation_name, e);
|
||||
E::from(e)
|
||||
})?;
|
||||
|
||||
|
||||
// Execute the operation within the transaction
|
||||
match operation(&mut tx).await {
|
||||
Ok(result) => {
|
||||
@@ -32,17 +34,20 @@ where
|
||||
Ok(_) => {
|
||||
debug!("Transaction committed successfully for: {}", operation_name);
|
||||
Ok(result)
|
||||
},
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Failed to commit transaction for {}: {}", operation_name, e);
|
||||
Err(E::from(e))
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
Err(e) => {
|
||||
// If operation fails, rollback the transaction
|
||||
if let Err(rollback_err) = tx.rollback().await {
|
||||
error!("Failed to rollback transaction for {}: {}", operation_name, rollback_err);
|
||||
error!(
|
||||
"Failed to rollback transaction for {}: {}",
|
||||
operation_name, rollback_err
|
||||
);
|
||||
// Still return the original error
|
||||
} else {
|
||||
info!("Transaction rolled back for {}: {}", operation_name, e);
|
||||
@@ -50,4 +55,4 @@ where
|
||||
Err(e)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,12 +1,14 @@
|
||||
use async_trait::async_trait;
|
||||
use futures::future::BoxFuture;
|
||||
use sqlx::{PgPool, Row};
|
||||
use std::sync::Arc;
|
||||
use futures::future::BoxFuture;
|
||||
|
||||
use crate::domain::entities::user::{User, UserRole};
|
||||
use crate::domain::repositories::user_repository::{UserRepository, UserRepositoryError, UserRepositoryResult, StorageStats};
|
||||
use crate::application::ports::auth_ports::UserStoragePort;
|
||||
use crate::common::errors::DomainError;
|
||||
use crate::domain::entities::user::{User, UserRole};
|
||||
use crate::domain::repositories::user_repository::{
|
||||
StorageStats, UserRepository, UserRepositoryError, UserRepositoryResult,
|
||||
};
|
||||
use crate::infrastructure::repositories::pg::transaction_utils::with_transaction;
|
||||
|
||||
// Implement From<sqlx::Error> for UserRepositoryError to allow automatic conversions
|
||||
@@ -24,28 +26,20 @@ impl UserPgRepository {
|
||||
pub fn new(pool: Arc<PgPool>) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
|
||||
// Helper method to map SQL errors to domain errors
|
||||
pub fn map_sqlx_error(err: sqlx::Error) -> UserRepositoryError {
|
||||
match err {
|
||||
sqlx::Error::RowNotFound => {
|
||||
UserRepositoryError::NotFound("User not found".to_string())
|
||||
},
|
||||
sqlx::Error::RowNotFound => UserRepositoryError::NotFound("User not found".to_string()),
|
||||
sqlx::Error::Database(db_err) => {
|
||||
if db_err.code().is_some_and(|code| code == "23505") {
|
||||
// PostgreSQL uniqueness violation code
|
||||
UserRepositoryError::AlreadyExists(
|
||||
"User or email already exists".to_string()
|
||||
)
|
||||
UserRepositoryError::AlreadyExists("User or email already exists".to_string())
|
||||
} else {
|
||||
UserRepositoryError::DatabaseError(
|
||||
format!("Database error: {}", db_err)
|
||||
)
|
||||
UserRepositoryError::DatabaseError(format!("Database error: {}", db_err))
|
||||
}
|
||||
},
|
||||
_ => UserRepositoryError::DatabaseError(
|
||||
format!("Database error: {}", err)
|
||||
),
|
||||
}
|
||||
_ => UserRepositoryError::DatabaseError(format!("Database error: {}", err)),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -56,21 +50,18 @@ impl UserRepository for UserPgRepository {
|
||||
async fn create_user(&self, user: User) -> UserRepositoryResult<User> {
|
||||
// Create a copy of the user for the closure
|
||||
let user_clone = user.clone();
|
||||
|
||||
with_transaction(
|
||||
&self.pool,
|
||||
"create_user",
|
||||
|tx| {
|
||||
// We need to move the closure into a BoxFuture to return inside
|
||||
// the with_transaction call
|
||||
Box::pin(async move {
|
||||
// Use getters to extract the values
|
||||
// Convert user.role() to string to pass it as plain text
|
||||
let role_str = user_clone.role().to_string();
|
||||
|
||||
// Modify the SQL to do an explicit cast to the auth.userrole type
|
||||
let _result = sqlx::query(
|
||||
r#"
|
||||
|
||||
with_transaction(&self.pool, "create_user", |tx| {
|
||||
// We need to move the closure into a BoxFuture to return inside
|
||||
// the with_transaction call
|
||||
Box::pin(async move {
|
||||
// Use getters to extract the values
|
||||
// Convert user.role() to string to pass it as plain text
|
||||
let role_str = user_clone.role().to_string();
|
||||
|
||||
// Modify the SQL to do an explicit cast to the auth.userrole type
|
||||
let _result = sqlx::query(
|
||||
r#"
|
||||
INSERT INTO auth.users (
|
||||
id, username, email, password_hash, role,
|
||||
storage_quota_bytes, storage_used_bytes,
|
||||
@@ -81,36 +72,36 @@ impl UserRepository for UserPgRepository {
|
||||
$12, $13
|
||||
)
|
||||
RETURNING *
|
||||
"#
|
||||
)
|
||||
.bind(user_clone.id())
|
||||
.bind(user_clone.username())
|
||||
.bind(user_clone.email())
|
||||
.bind(user_clone.password_hash())
|
||||
.bind(&role_str) // Convert to string but with explicit cast in SQL
|
||||
.bind(user_clone.storage_quota_bytes())
|
||||
.bind(user_clone.storage_used_bytes())
|
||||
.bind(user_clone.created_at())
|
||||
.bind(user_clone.updated_at())
|
||||
.bind(user_clone.last_login_at())
|
||||
.bind(user_clone.is_active())
|
||||
.bind(user_clone.oidc_provider())
|
||||
.bind(user_clone.oidc_subject())
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
// We could perform additional operations here,
|
||||
// such as configuring permissions, roles, etc.
|
||||
|
||||
Ok(user_clone)
|
||||
}) as BoxFuture<'_, UserRepositoryResult<User>>
|
||||
}
|
||||
).await?;
|
||||
|
||||
"#,
|
||||
)
|
||||
.bind(user_clone.id())
|
||||
.bind(user_clone.username())
|
||||
.bind(user_clone.email())
|
||||
.bind(user_clone.password_hash())
|
||||
.bind(&role_str) // Convert to string but with explicit cast in SQL
|
||||
.bind(user_clone.storage_quota_bytes())
|
||||
.bind(user_clone.storage_used_bytes())
|
||||
.bind(user_clone.created_at())
|
||||
.bind(user_clone.updated_at())
|
||||
.bind(user_clone.last_login_at())
|
||||
.bind(user_clone.is_active())
|
||||
.bind(user_clone.oidc_provider())
|
||||
.bind(user_clone.oidc_subject())
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
// We could perform additional operations here,
|
||||
// such as configuring permissions, roles, etc.
|
||||
|
||||
Ok(user_clone)
|
||||
}) as BoxFuture<'_, UserRepositoryResult<User>>
|
||||
})
|
||||
.await?;
|
||||
|
||||
Ok(user) // Return the original user for simplicity
|
||||
}
|
||||
|
||||
|
||||
/// Gets a user by ID
|
||||
async fn get_user_by_id(&self, id: &str) -> UserRepositoryResult<User> {
|
||||
let row = sqlx::query(
|
||||
@@ -122,7 +113,7 @@ impl UserRepository for UserPgRepository {
|
||||
oidc_provider, oidc_subject
|
||||
FROM auth.users
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_one(&*self.pool)
|
||||
@@ -135,7 +126,7 @@ impl UserRepository for UserPgRepository {
|
||||
Some("admin") => UserRole::Admin,
|
||||
_ => UserRole::User,
|
||||
};
|
||||
|
||||
|
||||
Ok(User::from_data_full(
|
||||
row.get("id"),
|
||||
row.get("username"),
|
||||
@@ -152,7 +143,7 @@ impl UserRepository for UserPgRepository {
|
||||
row.get("oidc_subject"),
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
/// Gets a user by username
|
||||
async fn get_user_by_username(&self, username: &str) -> UserRepositoryResult<User> {
|
||||
let row = sqlx::query(
|
||||
@@ -164,7 +155,7 @@ impl UserRepository for UserPgRepository {
|
||||
oidc_provider, oidc_subject
|
||||
FROM auth.users
|
||||
WHERE username = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(username)
|
||||
.fetch_one(&*self.pool)
|
||||
@@ -177,7 +168,7 @@ impl UserRepository for UserPgRepository {
|
||||
Some("admin") => UserRole::Admin,
|
||||
_ => UserRole::User,
|
||||
};
|
||||
|
||||
|
||||
Ok(User::from_data_full(
|
||||
row.get("id"),
|
||||
row.get("username"),
|
||||
@@ -194,7 +185,7 @@ impl UserRepository for UserPgRepository {
|
||||
row.get("oidc_subject"),
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
/// Gets a user by email
|
||||
async fn get_user_by_email(&self, email: &str) -> UserRepositoryResult<User> {
|
||||
let row = sqlx::query(
|
||||
@@ -206,7 +197,7 @@ impl UserRepository for UserPgRepository {
|
||||
oidc_provider, oidc_subject
|
||||
FROM auth.users
|
||||
WHERE email = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(email)
|
||||
.fetch_one(&*self.pool)
|
||||
@@ -219,7 +210,7 @@ impl UserRepository for UserPgRepository {
|
||||
Some("admin") => UserRole::Admin,
|
||||
_ => UserRole::User,
|
||||
};
|
||||
|
||||
|
||||
Ok(User::from_data_full(
|
||||
row.get("id"),
|
||||
row.get("username"),
|
||||
@@ -236,20 +227,17 @@ impl UserRepository for UserPgRepository {
|
||||
row.get("oidc_subject"),
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
/// Updates an existing user using a transaction
|
||||
async fn update_user(&self, user: User) -> UserRepositoryResult<User> {
|
||||
// Create a copy of the user for the closure
|
||||
let user_clone = user.clone();
|
||||
|
||||
with_transaction(
|
||||
&self.pool,
|
||||
"update_user",
|
||||
|tx| {
|
||||
Box::pin(async move {
|
||||
// Update the user
|
||||
sqlx::query(
|
||||
r#"
|
||||
|
||||
with_transaction(&self.pool, "update_user", |tx| {
|
||||
Box::pin(async move {
|
||||
// Update the user
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE auth.users
|
||||
SET
|
||||
username = $2,
|
||||
@@ -262,35 +250,39 @@ impl UserRepository for UserPgRepository {
|
||||
last_login_at = $9,
|
||||
active = $10
|
||||
WHERE id = $1
|
||||
"#
|
||||
)
|
||||
.bind(user_clone.id())
|
||||
.bind(user_clone.username())
|
||||
.bind(user_clone.email())
|
||||
.bind(user_clone.password_hash())
|
||||
.bind(user_clone.role().to_string())
|
||||
.bind(user_clone.storage_quota_bytes())
|
||||
.bind(user_clone.storage_used_bytes())
|
||||
.bind(user_clone.updated_at())
|
||||
.bind(user_clone.last_login_at())
|
||||
.bind(user_clone.is_active())
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
// We could perform additional operations here inside
|
||||
// the same transaction, such as updating permissions, etc.
|
||||
|
||||
Ok(user_clone)
|
||||
}) as BoxFuture<'_, UserRepositoryResult<User>>
|
||||
}
|
||||
).await?;
|
||||
|
||||
"#,
|
||||
)
|
||||
.bind(user_clone.id())
|
||||
.bind(user_clone.username())
|
||||
.bind(user_clone.email())
|
||||
.bind(user_clone.password_hash())
|
||||
.bind(user_clone.role().to_string())
|
||||
.bind(user_clone.storage_quota_bytes())
|
||||
.bind(user_clone.storage_used_bytes())
|
||||
.bind(user_clone.updated_at())
|
||||
.bind(user_clone.last_login_at())
|
||||
.bind(user_clone.is_active())
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
// We could perform additional operations here inside
|
||||
// the same transaction, such as updating permissions, etc.
|
||||
|
||||
Ok(user_clone)
|
||||
}) as BoxFuture<'_, UserRepositoryResult<User>>
|
||||
})
|
||||
.await?;
|
||||
|
||||
Ok(user)
|
||||
}
|
||||
|
||||
|
||||
/// Updates only the storage usage of a user
|
||||
async fn update_storage_usage(&self, user_id: &str, usage_bytes: i64) -> UserRepositoryResult<()> {
|
||||
async fn update_storage_usage(
|
||||
&self,
|
||||
user_id: &str,
|
||||
usage_bytes: i64,
|
||||
) -> UserRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE auth.users
|
||||
@@ -298,7 +290,7 @@ impl UserRepository for UserPgRepository {
|
||||
storage_used_bytes = $2,
|
||||
updated_at = NOW()
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.bind(usage_bytes)
|
||||
@@ -308,7 +300,7 @@ impl UserRepository for UserPgRepository {
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
/// Updates the last login date
|
||||
async fn update_last_login(&self, user_id: &str) -> UserRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
@@ -318,7 +310,7 @@ impl UserRepository for UserPgRepository {
|
||||
last_login_at = NOW(),
|
||||
updated_at = NOW()
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.execute(&*self.pool)
|
||||
@@ -327,7 +319,7 @@ impl UserRepository for UserPgRepository {
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
/// Lists users with pagination
|
||||
async fn list_users(&self, limit: i64, offset: i64) -> UserRepositoryResult<Vec<User>> {
|
||||
let rows = sqlx::query(
|
||||
@@ -340,7 +332,7 @@ impl UserRepository for UserPgRepository {
|
||||
FROM auth.users
|
||||
ORDER BY created_at DESC
|
||||
LIMIT $1 OFFSET $2
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(limit)
|
||||
.bind(offset)
|
||||
@@ -348,7 +340,8 @@ impl UserRepository for UserPgRepository {
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
let users = rows.into_iter()
|
||||
let users = rows
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
// Convert role string to UserRole enum for each row
|
||||
let role_str: Option<String> = row.try_get("role_text").unwrap_or(None);
|
||||
@@ -356,7 +349,7 @@ impl UserRepository for UserPgRepository {
|
||||
Some("admin") => UserRole::Admin,
|
||||
_ => UserRole::User,
|
||||
};
|
||||
|
||||
|
||||
User::from_data_full(
|
||||
row.get("id"),
|
||||
row.get("username"),
|
||||
@@ -377,9 +370,13 @@ impl UserRepository for UserPgRepository {
|
||||
|
||||
Ok(users)
|
||||
}
|
||||
|
||||
|
||||
/// Activates or deactivates a user
|
||||
async fn set_user_active_status(&self, user_id: &str, active: bool) -> UserRepositoryResult<()> {
|
||||
async fn set_user_active_status(
|
||||
&self,
|
||||
user_id: &str,
|
||||
active: bool,
|
||||
) -> UserRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE auth.users
|
||||
@@ -387,7 +384,7 @@ impl UserRepository for UserPgRepository {
|
||||
active = $2,
|
||||
updated_at = NOW()
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.bind(active)
|
||||
@@ -397,9 +394,13 @@ impl UserRepository for UserPgRepository {
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
/// Changes a user's password
|
||||
async fn change_password(&self, user_id: &str, password_hash: &str) -> UserRepositoryResult<()> {
|
||||
async fn change_password(
|
||||
&self,
|
||||
user_id: &str,
|
||||
password_hash: &str,
|
||||
) -> UserRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE auth.users
|
||||
@@ -407,7 +408,7 @@ impl UserRepository for UserPgRepository {
|
||||
password_hash = $2,
|
||||
updated_at = NOW()
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.bind(password_hash)
|
||||
@@ -417,12 +418,12 @@ impl UserRepository for UserPgRepository {
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
/// Changes a user's role
|
||||
async fn change_role(&self, user_id: &str, role: UserRole) -> UserRepositoryResult<()> {
|
||||
// Convert the role to string for the binding
|
||||
let role_str = role.to_string();
|
||||
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE auth.users
|
||||
@@ -430,7 +431,7 @@ impl UserRepository for UserPgRepository {
|
||||
role = $2::auth.userrole,
|
||||
updated_at = NOW()
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.bind(&role_str)
|
||||
@@ -440,7 +441,7 @@ impl UserRepository for UserPgRepository {
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
/// Lists users by role
|
||||
async fn list_users_by_role(&self, role: &str) -> UserRepositoryResult<Vec<User>> {
|
||||
let rows = sqlx::query(
|
||||
@@ -453,14 +454,15 @@ impl UserRepository for UserPgRepository {
|
||||
FROM auth.users
|
||||
WHERE role::text = $1
|
||||
ORDER BY created_at DESC
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(role)
|
||||
.fetch_all(&*self.pool)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
let users = rows.into_iter()
|
||||
let users = rows
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
// Convert role string to UserRole enum for each row
|
||||
let role_str: Option<String> = row.try_get("role_text").unwrap_or(None);
|
||||
@@ -468,7 +470,7 @@ impl UserRepository for UserPgRepository {
|
||||
Some("admin") => UserRole::Admin,
|
||||
_ => UserRole::User,
|
||||
};
|
||||
|
||||
|
||||
User::from_data_full(
|
||||
row.get("id"),
|
||||
row.get("username"),
|
||||
@@ -489,14 +491,14 @@ impl UserRepository for UserPgRepository {
|
||||
|
||||
Ok(users)
|
||||
}
|
||||
|
||||
|
||||
/// Deletes a user
|
||||
async fn delete_user(&self, user_id: &str) -> UserRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
DELETE FROM auth.users
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.execute(&*self.pool)
|
||||
@@ -507,7 +509,11 @@ impl UserRepository for UserPgRepository {
|
||||
}
|
||||
|
||||
/// Finds a user by OIDC provider + subject pair
|
||||
async fn get_user_by_oidc_subject(&self, provider: &str, subject: &str) -> UserRepositoryResult<User> {
|
||||
async fn get_user_by_oidc_subject(
|
||||
&self,
|
||||
provider: &str,
|
||||
subject: &str,
|
||||
) -> UserRepositoryResult<User> {
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -517,7 +523,7 @@ impl UserRepository for UserPgRepository {
|
||||
oidc_provider, oidc_subject
|
||||
FROM auth.users
|
||||
WHERE oidc_provider = $1 AND oidc_subject = $2
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(provider)
|
||||
.bind(subject)
|
||||
@@ -549,7 +555,11 @@ impl UserRepository for UserPgRepository {
|
||||
}
|
||||
|
||||
/// Updates a user's storage quota
|
||||
async fn update_storage_quota(&self, user_id: &str, quota_bytes: i64) -> UserRepositoryResult<()> {
|
||||
async fn update_storage_quota(
|
||||
&self,
|
||||
user_id: &str,
|
||||
quota_bytes: i64,
|
||||
) -> UserRepositoryResult<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE auth.users
|
||||
@@ -557,7 +567,7 @@ impl UserRepository for UserPgRepository {
|
||||
storage_quota_bytes = $2,
|
||||
updated_at = NOW()
|
||||
WHERE id = $1
|
||||
"#
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.bind(quota_bytes)
|
||||
@@ -570,12 +580,10 @@ impl UserRepository for UserPgRepository {
|
||||
|
||||
/// Counts the total number of users
|
||||
async fn count_users(&self) -> UserRepositoryResult<i64> {
|
||||
let row = sqlx::query(
|
||||
"SELECT COUNT(*) as count FROM auth.users"
|
||||
)
|
||||
.fetch_one(&*self.pool)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
let row = sqlx::query("SELECT COUNT(*) as count FROM auth.users")
|
||||
.fetch_one(&*self.pool)
|
||||
.await
|
||||
.map_err(Self::map_sqlx_error)?;
|
||||
|
||||
let count: i64 = row.get("count");
|
||||
Ok(count)
|
||||
@@ -614,52 +622,74 @@ impl UserRepository for UserPgRepository {
|
||||
#[async_trait]
|
||||
impl UserStoragePort for UserPgRepository {
|
||||
async fn create_user(&self, user: User) -> Result<User, DomainError> {
|
||||
UserRepository::create_user(self, user).await.map_err(DomainError::from)
|
||||
UserRepository::create_user(self, user)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
|
||||
async fn get_user_by_id(&self, id: &str) -> Result<User, DomainError> {
|
||||
UserRepository::get_user_by_id(self, id).await.map_err(DomainError::from)
|
||||
UserRepository::get_user_by_id(self, id)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
|
||||
async fn get_user_by_username(&self, username: &str) -> Result<User, DomainError> {
|
||||
UserRepository::get_user_by_username(self, username).await.map_err(DomainError::from)
|
||||
UserRepository::get_user_by_username(self, username)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
|
||||
async fn get_user_by_email(&self, email: &str) -> Result<User, DomainError> {
|
||||
UserRepository::get_user_by_email(self, email).await.map_err(DomainError::from)
|
||||
UserRepository::get_user_by_email(self, email)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
|
||||
async fn update_user(&self, user: User) -> Result<User, DomainError> {
|
||||
UserRepository::update_user(self, user).await.map_err(DomainError::from)
|
||||
UserRepository::update_user(self, user)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
async fn update_storage_usage(&self, user_id: &str, usage_bytes: i64) -> Result<(), DomainError> {
|
||||
|
||||
async fn update_storage_usage(
|
||||
&self,
|
||||
user_id: &str,
|
||||
usage_bytes: i64,
|
||||
) -> Result<(), DomainError> {
|
||||
UserRepository::update_storage_usage(self, user_id, usage_bytes)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
|
||||
async fn list_users(&self, limit: i64, offset: i64) -> Result<Vec<User>, DomainError> {
|
||||
UserRepository::list_users(self, limit, offset).await.map_err(DomainError::from)
|
||||
UserRepository::list_users(self, limit, offset)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
|
||||
async fn list_users_by_role(&self, role: &str) -> Result<Vec<User>, DomainError> {
|
||||
UserRepository::list_users_by_role(self, role).await.map_err(DomainError::from)
|
||||
UserRepository::list_users_by_role(self, role)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
|
||||
async fn delete_user(&self, user_id: &str) -> Result<(), DomainError> {
|
||||
UserRepository::delete_user(self, user_id)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
|
||||
async fn change_password(&self, user_id: &str, password_hash: &str) -> Result<(), DomainError> {
|
||||
UserRepository::change_password(self, user_id, password_hash)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
async fn get_user_by_oidc_subject(&self, provider: &str, subject: &str) -> Result<User, DomainError> {
|
||||
async fn get_user_by_oidc_subject(
|
||||
&self,
|
||||
provider: &str,
|
||||
subject: &str,
|
||||
) -> Result<User, DomainError> {
|
||||
UserRepository::get_user_by_oidc_subject(self, provider, subject)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
@@ -681,7 +711,11 @@ impl UserStoragePort for UserPgRepository {
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
|
||||
async fn update_storage_quota(&self, user_id: &str, quota_bytes: i64) -> Result<(), DomainError> {
|
||||
async fn update_storage_quota(
|
||||
&self,
|
||||
user_id: &str,
|
||||
quota_bytes: i64,
|
||||
) -> Result<(), DomainError> {
|
||||
UserRepository::update_storage_quota(self, user_id, quota_bytes)
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
@@ -692,4 +726,4 @@ impl UserStoragePort for UserPgRepository {
|
||||
.await
|
||||
.map_err(DomainError::from)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,22 +12,22 @@ use crate::common::errors::DomainError;
|
||||
pub enum FileRepositoryError {
|
||||
#[error("File not found: {0}")]
|
||||
NotFound(String),
|
||||
|
||||
|
||||
#[error("File already exists: {0}")]
|
||||
AlreadyExists(String),
|
||||
|
||||
|
||||
#[error("Invalid file path: {0}")]
|
||||
InvalidPath(String),
|
||||
|
||||
|
||||
#[error("Operation not supported: {0}")]
|
||||
OperationNotSupported(String),
|
||||
|
||||
|
||||
#[error("Storage error: {0}")]
|
||||
StorageError(String),
|
||||
|
||||
|
||||
#[error("Domain error: {0}")]
|
||||
DomainError(#[from] DomainError),
|
||||
|
||||
|
||||
#[error("Other error: {0}")]
|
||||
Other(String),
|
||||
}
|
||||
@@ -39,25 +39,25 @@ pub type FileRepositoryResult<T> = Result<T, FileRepositoryError>;
|
||||
pub enum FolderRepositoryError {
|
||||
#[error("Folder not found: {0}")]
|
||||
NotFound(String),
|
||||
|
||||
|
||||
#[error("Folder already exists: {0}")]
|
||||
AlreadyExists(String),
|
||||
|
||||
|
||||
#[error("Invalid folder path: {0}")]
|
||||
InvalidPath(String),
|
||||
|
||||
|
||||
#[error("Operation not supported: {0}")]
|
||||
OperationNotSupported(String),
|
||||
|
||||
|
||||
#[error("Storage error: {0}")]
|
||||
StorageError(String),
|
||||
|
||||
|
||||
#[error("Validation error: {0}")]
|
||||
ValidationError(String),
|
||||
|
||||
|
||||
#[error("Domain error: {0}")]
|
||||
DomainError(#[from] DomainError),
|
||||
|
||||
|
||||
#[error("Other error: {0}")]
|
||||
Other(String),
|
||||
}
|
||||
@@ -71,10 +71,16 @@ impl From<FileRepositoryError> for DomainError {
|
||||
match err {
|
||||
FileRepositoryError::NotFound(id) => DomainError::not_found("File", id),
|
||||
FileRepositoryError::AlreadyExists(path) => DomainError::already_exists("File", path),
|
||||
FileRepositoryError::InvalidPath(path) => DomainError::validation_error(format!("Invalid path: {}", path)),
|
||||
FileRepositoryError::StorageError(msg) => DomainError::internal_error("File", format!("Storage error: {}", msg)),
|
||||
FileRepositoryError::InvalidPath(path) => {
|
||||
DomainError::validation_error(format!("Invalid path: {}", path))
|
||||
}
|
||||
FileRepositoryError::StorageError(msg) => {
|
||||
DomainError::internal_error("File", format!("Storage error: {}", msg))
|
||||
}
|
||||
FileRepositoryError::Other(msg) => DomainError::internal_error("File", msg),
|
||||
FileRepositoryError::OperationNotSupported(msg) => DomainError::operation_not_supported("File", msg),
|
||||
FileRepositoryError::OperationNotSupported(msg) => {
|
||||
DomainError::operation_not_supported("File", msg)
|
||||
}
|
||||
FileRepositoryError::DomainError(e) => e,
|
||||
}
|
||||
}
|
||||
@@ -84,12 +90,20 @@ impl From<FolderRepositoryError> for DomainError {
|
||||
fn from(err: FolderRepositoryError) -> Self {
|
||||
match err {
|
||||
FolderRepositoryError::NotFound(id) => DomainError::not_found("Folder", id),
|
||||
FolderRepositoryError::AlreadyExists(path) => DomainError::already_exists("Folder", path),
|
||||
FolderRepositoryError::InvalidPath(path) => DomainError::validation_error(format!("Invalid path: {}", path)),
|
||||
FolderRepositoryError::StorageError(msg) => DomainError::internal_error("Folder", format!("Storage error: {}", msg)),
|
||||
FolderRepositoryError::AlreadyExists(path) => {
|
||||
DomainError::already_exists("Folder", path)
|
||||
}
|
||||
FolderRepositoryError::InvalidPath(path) => {
|
||||
DomainError::validation_error(format!("Invalid path: {}", path))
|
||||
}
|
||||
FolderRepositoryError::StorageError(msg) => {
|
||||
DomainError::internal_error("Folder", format!("Storage error: {}", msg))
|
||||
}
|
||||
FolderRepositoryError::ValidationError(msg) => DomainError::validation_error(msg),
|
||||
FolderRepositoryError::Other(msg) => DomainError::internal_error("Folder", msg),
|
||||
FolderRepositoryError::OperationNotSupported(msg) => DomainError::operation_not_supported("Folder", msg),
|
||||
FolderRepositoryError::OperationNotSupported(msg) => {
|
||||
DomainError::operation_not_supported("Folder", msg)
|
||||
}
|
||||
FolderRepositoryError::DomainError(e) => e,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,9 +7,7 @@ use tokio::{fs, io};
|
||||
use crate::{
|
||||
application::ports::share_ports::ShareStoragePort,
|
||||
common::{config::AppConfig, errors::DomainError},
|
||||
domain::{
|
||||
entities::share::{Share, ShareItemType},
|
||||
},
|
||||
domain::entities::share::{Share, ShareItemType},
|
||||
};
|
||||
|
||||
// Structure for storing in the file system
|
||||
@@ -74,8 +72,8 @@ impl ShareFsRepository {
|
||||
|
||||
/// Converts a file system record to a domain entity
|
||||
fn to_entity(&self, record: &ShareRecord) -> Share {
|
||||
let item_type = ShareItemType::try_from(record.item_type.as_str())
|
||||
.unwrap_or(ShareItemType::File);
|
||||
let item_type =
|
||||
ShareItemType::try_from(record.item_type.as_str()).unwrap_or(ShareItemType::File);
|
||||
|
||||
let permissions = crate::domain::entities::share::SharePermissions::new(
|
||||
record.permissions_read,
|
||||
@@ -119,7 +117,9 @@ impl ShareFsRepository {
|
||||
#[async_trait]
|
||||
impl ShareStoragePort for ShareFsRepository {
|
||||
async fn save_share(&self, share: &Share) -> Result<Share, DomainError> {
|
||||
let mut shares = self.read_shares().await
|
||||
let mut shares = self
|
||||
.read_shares()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("Share", e.to_string()))?;
|
||||
|
||||
// Check if the link already exists
|
||||
@@ -135,21 +135,22 @@ impl ShareStoragePort for ShareFsRepository {
|
||||
shares.push(record);
|
||||
}
|
||||
|
||||
self.write_shares(&shares).await
|
||||
self.write_shares(&shares)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("Share", e.to_string()))?;
|
||||
|
||||
Ok(share.clone())
|
||||
}
|
||||
|
||||
async fn find_share_by_id(&self, id: &str) -> Result<Share, DomainError> {
|
||||
let shares = self.read_shares().await
|
||||
let shares = self
|
||||
.read_shares()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("Share", e.to_string()))?;
|
||||
|
||||
let share = shares.iter()
|
||||
.find(|s| s.id == id)
|
||||
.ok_or_else(|| {
|
||||
DomainError::not_found("Share", format!("Share with ID {} not found", id))
|
||||
});
|
||||
let share = shares.iter().find(|s| s.id == id).ok_or_else(|| {
|
||||
DomainError::not_found("Share", format!("Share with ID {} not found", id))
|
||||
});
|
||||
|
||||
match share {
|
||||
Ok(record) => Ok(self.to_entity(record)),
|
||||
@@ -158,14 +159,14 @@ impl ShareStoragePort for ShareFsRepository {
|
||||
}
|
||||
|
||||
async fn find_share_by_token(&self, token: &str) -> Result<Share, DomainError> {
|
||||
let shares = self.read_shares().await
|
||||
let shares = self
|
||||
.read_shares()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("Share", e.to_string()))?;
|
||||
|
||||
let share = shares.iter()
|
||||
.find(|s| s.token == token)
|
||||
.ok_or_else(|| {
|
||||
DomainError::not_found("Share", format!("Share with token {} not found", token))
|
||||
});
|
||||
let share = shares.iter().find(|s| s.token == token).ok_or_else(|| {
|
||||
DomainError::not_found("Share", format!("Share with token {} not found", token))
|
||||
});
|
||||
|
||||
match share {
|
||||
Ok(record) => Ok(self.to_entity(record)),
|
||||
@@ -173,12 +174,19 @@ impl ShareStoragePort for ShareFsRepository {
|
||||
}
|
||||
}
|
||||
|
||||
async fn find_shares_by_item(&self, item_id: &str, item_type: &ShareItemType) -> Result<Vec<Share>, DomainError> {
|
||||
let shares = self.read_shares().await
|
||||
async fn find_shares_by_item(
|
||||
&self,
|
||||
item_id: &str,
|
||||
item_type: &ShareItemType,
|
||||
) -> Result<Vec<Share>, DomainError> {
|
||||
let shares = self
|
||||
.read_shares()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("Share", e.to_string()))?;
|
||||
|
||||
let type_str = item_type.to_string();
|
||||
let result: Vec<Share> = shares.iter()
|
||||
let result: Vec<Share> = shares
|
||||
.iter()
|
||||
.filter(|s| s.item_id == item_id && s.item_type == type_str)
|
||||
.map(|record| self.to_entity(record))
|
||||
.collect();
|
||||
@@ -187,27 +195,37 @@ impl ShareStoragePort for ShareFsRepository {
|
||||
}
|
||||
|
||||
async fn update_share(&self, share: &Share) -> Result<Share, DomainError> {
|
||||
let mut shares = self.read_shares().await
|
||||
let mut shares = self
|
||||
.read_shares()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("Share", e.to_string()))?;
|
||||
|
||||
// Find the index of the link to update
|
||||
let index = shares.iter().position(|s| s.id == share.id())
|
||||
let index = shares
|
||||
.iter()
|
||||
.position(|s| s.id == share.id())
|
||||
.ok_or_else(|| {
|
||||
DomainError::not_found("Share", format!("Share with ID {} not found for update", share.id()))
|
||||
DomainError::not_found(
|
||||
"Share",
|
||||
format!("Share with ID {} not found for update", share.id()),
|
||||
)
|
||||
})?;
|
||||
|
||||
// Update the record
|
||||
shares[index] = self.to_record(share);
|
||||
|
||||
// Save changes
|
||||
self.write_shares(&shares).await
|
||||
self.write_shares(&shares)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("Share", e.to_string()))?;
|
||||
|
||||
Ok(share.clone())
|
||||
}
|
||||
|
||||
async fn delete_share(&self, id: &str) -> Result<(), DomainError> {
|
||||
let mut shares = self.read_shares().await
|
||||
let mut shares = self
|
||||
.read_shares()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("Share", e.to_string()))?;
|
||||
|
||||
// Find the index of the link to delete
|
||||
@@ -216,22 +234,34 @@ impl ShareStoragePort for ShareFsRepository {
|
||||
|
||||
// If no link was deleted, it means it didn't exist
|
||||
if shares.len() == initial_len {
|
||||
return Err(DomainError::not_found("Share", format!("Share with ID {} not found for deletion", id)));
|
||||
return Err(DomainError::not_found(
|
||||
"Share",
|
||||
format!("Share with ID {} not found for deletion", id),
|
||||
));
|
||||
}
|
||||
|
||||
// Save changes
|
||||
self.write_shares(&shares).await
|
||||
self.write_shares(&shares)
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("Share", e.to_string()))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn find_shares_by_user(&self, user_id: &str, offset: usize, limit: usize) -> Result<(Vec<Share>, usize), DomainError> {
|
||||
let shares = self.read_shares().await
|
||||
async fn find_shares_by_user(
|
||||
&self,
|
||||
user_id: &str,
|
||||
offset: usize,
|
||||
limit: usize,
|
||||
) -> Result<(Vec<Share>, usize), DomainError> {
|
||||
let shares = self
|
||||
.read_shares()
|
||||
.await
|
||||
.map_err(|e| DomainError::internal_error("Share", e.to_string()))?;
|
||||
|
||||
// Filter the user's links
|
||||
let user_shares: Vec<ShareRecord> = shares.into_iter()
|
||||
let user_shares: Vec<ShareRecord> = shares
|
||||
.into_iter()
|
||||
.filter(|s| s.created_by == user_id)
|
||||
.collect();
|
||||
|
||||
@@ -239,7 +269,8 @@ impl ShareStoragePort for ShareFsRepository {
|
||||
let total = user_shares.len();
|
||||
|
||||
// Apply pagination
|
||||
let paginated: Vec<Share> = user_shares.iter()
|
||||
let paginated: Vec<Share> = user_shares
|
||||
.iter()
|
||||
.skip(offset)
|
||||
.take(limit)
|
||||
.map(|record| self.to_entity(record))
|
||||
|
||||
@@ -1,16 +1,16 @@
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::Arc;
|
||||
use async_trait::async_trait;
|
||||
use chrono::Utc;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::Arc;
|
||||
use tokio::fs;
|
||||
use uuid::Uuid;
|
||||
use tracing::{debug, error, instrument};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::common::errors::{Result, DomainError, ErrorKind};
|
||||
use crate::application::ports::outbound::IdMappingPort;
|
||||
use crate::common::errors::{DomainError, ErrorKind, Result};
|
||||
use crate::domain::entities::trashed_item::{TrashedItem, TrashedItemType};
|
||||
use crate::domain::repositories::trash_repository::TrashRepository;
|
||||
use crate::application::ports::outbound::IdMappingPort;
|
||||
|
||||
/// Structure for storing trash items in JSON format
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
@@ -38,162 +38,200 @@ impl TrashFsRepository {
|
||||
) -> Self {
|
||||
let trash_dir = storage_root.as_ref().join(".trash");
|
||||
let trash_index_path = trash_dir.join("trash_index.json");
|
||||
|
||||
|
||||
Self {
|
||||
trash_dir,
|
||||
trash_index_path,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Ensures the trash directory exists
|
||||
async fn ensure_trash_dir(&self) -> Result<()> {
|
||||
debug!("Checking if trash directory exists: {}", self.trash_dir.display());
|
||||
debug!(
|
||||
"Checking if trash directory exists: {}",
|
||||
self.trash_dir.display()
|
||||
);
|
||||
if !self.trash_dir.exists() {
|
||||
debug!("Trash directory does not exist, creating it: {}", self.trash_dir.display());
|
||||
fs::create_dir_all(&self.trash_dir).await
|
||||
.map_err(|e| {
|
||||
error!("Failed to create trash directory {}: {}", self.trash_dir.display(), e);
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!("Failed to create trash directory {}: {}", self.trash_dir.display(), e)
|
||||
)
|
||||
})?;
|
||||
debug!(
|
||||
"Trash directory does not exist, creating it: {}",
|
||||
self.trash_dir.display()
|
||||
);
|
||||
fs::create_dir_all(&self.trash_dir).await.map_err(|e| {
|
||||
error!(
|
||||
"Failed to create trash directory {}: {}",
|
||||
self.trash_dir.display(),
|
||||
e
|
||||
);
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!(
|
||||
"Failed to create trash directory {}: {}",
|
||||
self.trash_dir.display(),
|
||||
e
|
||||
),
|
||||
)
|
||||
})?;
|
||||
debug!("Trash directory created successfully");
|
||||
} else {
|
||||
debug!("Trash directory already exists");
|
||||
}
|
||||
|
||||
|
||||
// Ensure the files directory exists
|
||||
let files_dir = self.trash_dir.join("files");
|
||||
debug!("Checking if trash files directory exists: {}", files_dir.display());
|
||||
debug!(
|
||||
"Checking if trash files directory exists: {}",
|
||||
files_dir.display()
|
||||
);
|
||||
if !files_dir.exists() {
|
||||
debug!("Trash files directory does not exist, creating it: {}", files_dir.display());
|
||||
fs::create_dir_all(&files_dir).await
|
||||
.map_err(|e| {
|
||||
error!("Failed to create trash files directory {}: {}", files_dir.display(), e);
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!("Failed to create trash files directory {}: {}", files_dir.display(), e)
|
||||
)
|
||||
})?;
|
||||
debug!(
|
||||
"Trash files directory does not exist, creating it: {}",
|
||||
files_dir.display()
|
||||
);
|
||||
fs::create_dir_all(&files_dir).await.map_err(|e| {
|
||||
error!(
|
||||
"Failed to create trash files directory {}: {}",
|
||||
files_dir.display(),
|
||||
e
|
||||
);
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!(
|
||||
"Failed to create trash files directory {}: {}",
|
||||
files_dir.display(),
|
||||
e
|
||||
),
|
||||
)
|
||||
})?;
|
||||
debug!("Trash files directory created successfully");
|
||||
} else {
|
||||
debug!("Trash files directory already exists");
|
||||
}
|
||||
|
||||
|
||||
// Also ensure the folders directory exists
|
||||
let folders_dir = self.trash_dir.join("folders");
|
||||
debug!("Checking if trash folders directory exists: {}", folders_dir.display());
|
||||
debug!(
|
||||
"Checking if trash folders directory exists: {}",
|
||||
folders_dir.display()
|
||||
);
|
||||
if !folders_dir.exists() {
|
||||
debug!("Trash folders directory does not exist, creating it: {}", folders_dir.display());
|
||||
fs::create_dir_all(&folders_dir).await
|
||||
.map_err(|e| {
|
||||
error!("Failed to create trash folders directory {}: {}", folders_dir.display(), e);
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!("Failed to create trash folders directory {}: {}", folders_dir.display(), e)
|
||||
)
|
||||
})?;
|
||||
debug!(
|
||||
"Trash folders directory does not exist, creating it: {}",
|
||||
folders_dir.display()
|
||||
);
|
||||
fs::create_dir_all(&folders_dir).await.map_err(|e| {
|
||||
error!(
|
||||
"Failed to create trash folders directory {}: {}",
|
||||
folders_dir.display(),
|
||||
e
|
||||
);
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!(
|
||||
"Failed to create trash folders directory {}: {}",
|
||||
folders_dir.display(),
|
||||
e
|
||||
),
|
||||
)
|
||||
})?;
|
||||
debug!("Trash folders directory created successfully");
|
||||
} else {
|
||||
debug!("Trash folders directory already exists");
|
||||
}
|
||||
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
/// Gets all entries from the trash index
|
||||
async fn get_trash_entries(&self) -> Result<Vec<TrashedItemEntry>> {
|
||||
self.ensure_trash_dir().await?;
|
||||
|
||||
|
||||
if !self.trash_index_path.exists() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let content = fs::read_to_string(&self.trash_index_path).await
|
||||
.map_err(|e| DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!("Failed to read trash index: {}", e)
|
||||
))?;
|
||||
|
||||
|
||||
let content = fs::read_to_string(&self.trash_index_path)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!("Failed to read trash index: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
if content.trim().is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let entries: Vec<TrashedItemEntry> = serde_json::from_str(&content)
|
||||
.map_err(|e| DomainError::new(
|
||||
|
||||
let entries: Vec<TrashedItemEntry> = serde_json::from_str(&content).map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!("Failed to parse trash index: {}", e)
|
||||
))?;
|
||||
|
||||
format!("Failed to parse trash index: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(entries)
|
||||
}
|
||||
|
||||
|
||||
/// Saves all entries to the trash index
|
||||
async fn save_trash_entries(&self, entries: Vec<TrashedItemEntry>) -> Result<()> {
|
||||
self.ensure_trash_dir().await?;
|
||||
|
||||
let json = serde_json::to_string_pretty(&entries)
|
||||
.map_err(|e| DomainError::new(
|
||||
|
||||
let json = serde_json::to_string_pretty(&entries).map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!("Failed to serialize trash index: {}", e)
|
||||
))?;
|
||||
|
||||
fs::write(&self.trash_index_path, json).await
|
||||
.map_err(|e| DomainError::new(
|
||||
format!("Failed to serialize trash index: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
fs::write(&self.trash_index_path, json).await.map_err(|e| {
|
||||
DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!("Failed to write trash index: {}", e)
|
||||
))?;
|
||||
|
||||
format!("Failed to write trash index: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
/// Converts a JSON entry to a TrashedItem entity
|
||||
fn entry_to_trashed_item(&self, entry: TrashedItemEntry) -> Result<TrashedItem> {
|
||||
let item_type = match entry.item_type.as_str() {
|
||||
"file" => TrashedItemType::File,
|
||||
"folder" => TrashedItemType::Folder,
|
||||
_ => return Err(DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Trash",
|
||||
format!("Invalid trashed item type: {}", entry.item_type)
|
||||
)),
|
||||
_ => {
|
||||
return Err(DomainError::new(
|
||||
ErrorKind::InvalidInput,
|
||||
"Trash",
|
||||
format!("Invalid trashed item type: {}", entry.item_type),
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
let original_id = Uuid::parse_str(&entry.original_id)
|
||||
.map_err(|e| DomainError::validation_error(
|
||||
format!("Invalid original ID format: {}", e)
|
||||
))?;
|
||||
|
||||
|
||||
let original_id = Uuid::parse_str(&entry.original_id).map_err(|e| {
|
||||
DomainError::validation_error(format!("Invalid original ID format: {}", e))
|
||||
})?;
|
||||
|
||||
let id = Uuid::parse_str(&entry.id)
|
||||
.map_err(|e| DomainError::validation_error(
|
||||
format!("Invalid ID format: {}", e)
|
||||
))?;
|
||||
|
||||
.map_err(|e| DomainError::validation_error(format!("Invalid ID format: {}", e)))?;
|
||||
|
||||
let user_id = Uuid::parse_str(&entry.user_id)
|
||||
.map_err(|e| DomainError::validation_error(
|
||||
format!("Invalid user ID format: {}", e)
|
||||
))?;
|
||||
|
||||
.map_err(|e| DomainError::validation_error(format!("Invalid user ID format: {}", e)))?;
|
||||
|
||||
let trashed_at = chrono::DateTime::parse_from_rfc3339(&entry.trashed_at)
|
||||
.map_err(|e| DomainError::validation_error(
|
||||
format!("Invalid trashed_at date: {}", e)
|
||||
))?
|
||||
.map_err(|e| DomainError::validation_error(format!("Invalid trashed_at date: {}", e)))?
|
||||
.with_timezone(&Utc);
|
||||
|
||||
|
||||
let deletion_date = chrono::DateTime::parse_from_rfc3339(&entry.deletion_date)
|
||||
.map_err(|e| DomainError::validation_error(
|
||||
format!("Invalid deletion_date: {}", e)
|
||||
))?
|
||||
.map_err(|e| DomainError::validation_error(format!("Invalid deletion_date: {}", e)))?
|
||||
.with_timezone(&Utc);
|
||||
|
||||
|
||||
Ok(TrashedItem::from_raw(
|
||||
id,
|
||||
original_id,
|
||||
@@ -205,7 +243,7 @@ impl TrashFsRepository {
|
||||
deletion_date,
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
/// Converts a TrashedItem entity to a JSON entry
|
||||
fn trashed_item_to_entry(&self, item: &TrashedItem) -> TrashedItemEntry {
|
||||
TrashedItemEntry {
|
||||
@@ -228,55 +266,72 @@ impl TrashFsRepository {
|
||||
impl TrashRepository for TrashFsRepository {
|
||||
#[instrument(skip(self))]
|
||||
async fn add_to_trash(&self, item: &TrashedItem) -> Result<()> {
|
||||
debug!("Adding item to trash: id={}, user={}", item.id(), item.user_id());
|
||||
|
||||
debug!(
|
||||
"Adding item to trash: id={}, user={}",
|
||||
item.id(),
|
||||
item.user_id()
|
||||
);
|
||||
|
||||
// Ensure the trash directory exists for this user
|
||||
let user_trash_dir = self.trash_dir.join("files").join(item.user_id().to_string());
|
||||
let user_trash_dir = self
|
||||
.trash_dir
|
||||
.join("files")
|
||||
.join(item.user_id().to_string());
|
||||
debug!("User trash directory path: {}", user_trash_dir.display());
|
||||
|
||||
|
||||
// Create the user-specific trash directory
|
||||
debug!("Creating user trash directory: {}", user_trash_dir.display());
|
||||
debug!(
|
||||
"Creating user trash directory: {}",
|
||||
user_trash_dir.display()
|
||||
);
|
||||
match fs::create_dir_all(&user_trash_dir).await {
|
||||
Ok(_) => debug!("User trash directory created successfully"),
|
||||
Err(e) => {
|
||||
error!("Failed to create user trash directory {}: {}", user_trash_dir.display(), e);
|
||||
error!(
|
||||
"Failed to create user trash directory {}: {}",
|
||||
user_trash_dir.display(),
|
||||
e
|
||||
);
|
||||
return Err(DomainError::new(
|
||||
ErrorKind::InternalError,
|
||||
"Trash",
|
||||
format!("Failed to create user trash directory: {}", e)
|
||||
format!("Failed to create user trash directory: {}", e),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Log the current trash entries before adding the new one
|
||||
let mut entries = self.get_trash_entries().await?;
|
||||
debug!("Current trash entries count: {}", entries.len());
|
||||
|
||||
|
||||
// Create the entry for the trash index
|
||||
let entry = self.trashed_item_to_entry(item);
|
||||
debug!("Created trash entry: id={}, original_id={}, name={}",
|
||||
entry.id, entry.original_id, entry.name);
|
||||
|
||||
debug!(
|
||||
"Created trash entry: id={}, original_id={}, name={}",
|
||||
entry.id, entry.original_id, entry.name
|
||||
);
|
||||
|
||||
// Add the entry to the index and save
|
||||
entries.push(entry);
|
||||
debug!("Saving updated trash index with {} entries", entries.len());
|
||||
self.save_trash_entries(entries).await?;
|
||||
debug!("Trash index updated successfully");
|
||||
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[instrument(skip(self))]
|
||||
async fn get_trash_items(&self, user_id: &Uuid) -> Result<Vec<TrashedItem>> {
|
||||
debug!("Getting trash items for user: {}", user_id);
|
||||
|
||||
|
||||
let entries = self.get_trash_entries().await?;
|
||||
|
||||
|
||||
let user_id_str = user_id.to_string();
|
||||
let user_entries = entries.into_iter()
|
||||
let user_entries = entries
|
||||
.into_iter()
|
||||
.filter(|entry| entry.user_id == user_id_str)
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
|
||||
let mut items = Vec::new();
|
||||
for entry in user_entries {
|
||||
match self.entry_to_trashed_item(entry) {
|
||||
@@ -284,27 +339,28 @@ impl TrashRepository for TrashFsRepository {
|
||||
Err(e) => error!("Error converting trash entry to item: {}", e),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Ok(items)
|
||||
}
|
||||
|
||||
#[instrument(skip(self))]
|
||||
async fn get_trash_item(&self, id: &Uuid, user_id: &Uuid) -> Result<Option<TrashedItem>> {
|
||||
debug!("Looking for item in trash: id={}, user={}", id, user_id);
|
||||
|
||||
|
||||
let entries = self.get_trash_entries().await?;
|
||||
|
||||
|
||||
let id_str = id.to_string();
|
||||
let user_id_str = user_id.to_string();
|
||||
|
||||
let item_entry = entries.into_iter()
|
||||
|
||||
let item_entry = entries
|
||||
.into_iter()
|
||||
.find(|entry| entry.id == id_str && entry.user_id == user_id_str);
|
||||
|
||||
|
||||
match item_entry {
|
||||
Some(entry) => {
|
||||
let item = self.entry_to_trashed_item(entry)?;
|
||||
Ok(Some(item))
|
||||
},
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
@@ -312,16 +368,16 @@ impl TrashRepository for TrashFsRepository {
|
||||
#[instrument(skip(self))]
|
||||
async fn restore_from_trash(&self, id: &Uuid, user_id: &Uuid) -> Result<()> {
|
||||
debug!("Restoring item from trash: id={}, user={}", id, user_id);
|
||||
|
||||
|
||||
let mut entries = self.get_trash_entries().await?;
|
||||
|
||||
|
||||
let id_str = id.to_string();
|
||||
let user_id_str = user_id.to_string();
|
||||
|
||||
let index = entries.iter().position(|entry|
|
||||
entry.id == id_str && entry.user_id == user_id_str
|
||||
);
|
||||
|
||||
|
||||
let index = entries
|
||||
.iter()
|
||||
.position(|entry| entry.id == id_str && entry.user_id == user_id_str);
|
||||
|
||||
if let Some(index) = index {
|
||||
entries.remove(index);
|
||||
self.save_trash_entries(entries).await?;
|
||||
@@ -333,8 +389,11 @@ impl TrashRepository for TrashFsRepository {
|
||||
|
||||
#[instrument(skip(self))]
|
||||
async fn delete_permanently(&self, id: &Uuid, user_id: &Uuid) -> Result<()> {
|
||||
debug!("Permanently deleting item from trash: id={}, user={}", id, user_id);
|
||||
|
||||
debug!(
|
||||
"Permanently deleting item from trash: id={}, user={}",
|
||||
id, user_id
|
||||
);
|
||||
|
||||
// Simply remove the entry from the index
|
||||
// Physical files will be deleted through the corresponding repository
|
||||
self.restore_from_trash(id, user_id).await
|
||||
@@ -343,25 +402,25 @@ impl TrashRepository for TrashFsRepository {
|
||||
#[instrument(skip(self))]
|
||||
async fn clear_trash(&self, user_id: &Uuid) -> Result<()> {
|
||||
debug!("Clearing trash for user: {}", user_id);
|
||||
|
||||
|
||||
let mut entries = self.get_trash_entries().await?;
|
||||
let user_id_str = user_id.to_string();
|
||||
|
||||
|
||||
entries.retain(|entry| entry.user_id != user_id_str);
|
||||
self.save_trash_entries(entries).await?;
|
||||
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[instrument(skip(self))]
|
||||
async fn get_expired_items(&self) -> Result<Vec<TrashedItem>> {
|
||||
debug!("Looking for expired trash items");
|
||||
|
||||
|
||||
let entries = self.get_trash_entries().await?;
|
||||
let now = Utc::now();
|
||||
|
||||
|
||||
let mut expired_items = Vec::new();
|
||||
|
||||
|
||||
for entry in entries {
|
||||
match chrono::DateTime::parse_from_rfc3339(&entry.deletion_date) {
|
||||
Ok(date) => {
|
||||
@@ -372,11 +431,11 @@ impl TrashRepository for TrashFsRepository {
|
||||
Err(e) => error!("Error converting expired trash entry: {}", e),
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
Err(e) => error!("Invalid date format in trash entry: {}", e),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Ok(expired_items)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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, ¬_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, ¬_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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user