perf: findings 6.1, 6.2, 2.6 — async Argon2, moka cache, full streaming migration

- 6.1: PasswordHasherPort now async_trait with spawn_blocking for Argon2
- 6.2: OIDC pending maps migrated from std::sync::Mutex to moka::sync::Cache with TTL
- 2.6: All file download paths migrated to 64KB streaming (get_file_stream / read_blob_stream)
  - WOPI, dedup, batch ZIP, file_retrieval_service consumers migrated
  - WebDAV COPY uses zero-copy dedup (copy_file)
  - Removed dead code: get_file_content, get_file_mmap, read_blob, read_blob_bytes
    from traits, impls, stubs, and mocks (18 files touched)
This commit is contained in:
Diocrafts
2026-02-23 00:51:46 +01:00
parent b501c4052b
commit 85908311dc
18 changed files with 235 additions and 305 deletions
@@ -10,29 +10,23 @@ use crate::common::config::OidcConfig;
use crate::common::errors::{DomainError, ErrorKind};
use crate::domain::entities::session::Session;
use crate::domain::entities::user::{User, UserRole};
use std::collections::HashMap;
use moka::sync::Cache;
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::RwLock;
use std::time::Instant;
/// Maximum age for pending OIDC flows (10 minutes)
const OIDC_FLOW_TTL_SECS: u64 = 600;
/// Maximum age for pending one-time token codes (60 seconds)
const OIDC_TOKEN_TTL_SECS: u64 = 60;
use std::time::Duration;
/// Tracks a pending OIDC authorization flow (CSRF + PKCE + nonce)
#[derive(Clone)]
struct PendingOidcFlow {
created_at: Instant,
pkce_verifier: String,
nonce: String,
}
/// Tracks a pending one-time token exchange after successful OIDC callback
#[derive(Clone)]
struct PendingOidcToken {
auth_response: AuthResponseDto,
created_at: Instant,
}
/// Interior state for OIDC — protected by RwLock for hot-reload.
@@ -54,10 +48,12 @@ pub struct AuthApplicationService {
/// Path to the storage directory, used for disk-space–aware quota calculation
storage_path: PathBuf,
oidc: RwLock<OidcState>,
/// Pending OIDC authorization flows keyed by state token (CSRF + PKCE + nonce)
pending_oidc_flows: Mutex<HashMap<String, PendingOidcFlow>>,
/// Pending one-time token codes for secure token delivery after OIDC callback
pending_oidc_tokens: Mutex<HashMap<String, PendingOidcToken>>,
/// Pending OIDC authorization flows keyed by state token (CSRF + PKCE + nonce).
/// Auto-expires after 10 minutes via moka TTL; max 10 000 entries for DoS protection.
pending_oidc_flows: Cache<String, PendingOidcFlow>,
/// Pending one-time token codes for secure token delivery after OIDC callback.
/// Auto-expires after 60 seconds via moka TTL; max 10 000 entries for DoS protection.
pending_oidc_tokens: Cache<String, PendingOidcToken>,
}
impl AuthApplicationService {
@@ -79,8 +75,14 @@ impl AuthApplicationService {
service: None,
config: None,
}),
pending_oidc_flows: Mutex::new(HashMap::new()),
pending_oidc_tokens: Mutex::new(HashMap::new()),
pending_oidc_flows: Cache::builder()
.max_capacity(10_000)
.time_to_live(Duration::from_secs(600))
.build(),
pending_oidc_tokens: Cache::builder()
.max_capacity(10_000)
.time_to_live(Duration::from_secs(60))
.build(),
}
}
@@ -300,7 +302,7 @@ impl AuthApplicationService {
}
// Hash the password using the infrastructure service
let password_hash = self.password_hasher.hash_password(&dto.password)?;
let password_hash = self.password_hasher.hash_password(&dto.password).await?;
// Create user with the pre-generated hash
let user = User::new(dto.username.clone(), dto.email, password_hash, role, quota).map_err(
@@ -346,7 +348,8 @@ impl AuthApplicationService {
// Verify password using the injected hasher
let is_valid = self
.password_hasher
.verify_password(&dto.password, user.password_hash())?;
.verify_password(&dto.password, user.password_hash())
.await?;
if !is_valid {
return Err(DomainError::new(
@@ -502,7 +505,8 @@ impl AuthApplicationService {
// Verify current password using the injected hasher
let is_valid = self
.password_hasher
.verify_password(&dto.current_password, user.password_hash())?;
.verify_password(&dto.current_password, user.password_hash())
.await?;
if !is_valid {
return Err(DomainError::new(
@@ -522,7 +526,7 @@ impl AuthApplicationService {
}
// Hash new password and update user
let new_hash = self.password_hasher.hash_password(&dto.new_password)?;
let new_hash = self.password_hasher.hash_password(&dto.new_password).await?;
user.update_password_hash(new_hash);
// Save updated user
@@ -641,13 +645,7 @@ impl AuthApplicationService {
let password_hash = self
.password_hasher
.hash_password(&dto.password)
.map_err(|e| {
DomainError::new(
ErrorKind::InternalError,
"User",
format!("Error hashing password: {}", e),
)
})?;
.await?;
// Create the new admin user
let user = User::new(
@@ -753,7 +751,7 @@ impl AuthApplicationService {
});
// Hash password
let password_hash = self.password_hasher.hash_password(&dto.password)?;
let password_hash = self.password_hasher.hash_password(&dto.password).await?;
// Create domain entity
let user =
@@ -806,7 +804,7 @@ impl AuthApplicationService {
"Password must be at least 8 characters long".to_string(),
));
}
let hash = self.password_hasher.hash_password(new_password)?;
let hash = self.password_hasher.hash_password(new_password).await?;
self.user_storage.change_password(user_id, &hash).await
}
@@ -917,22 +915,14 @@ impl AuthApplicationService {
base64_url_encode(&hash)
};
// Store pending flow
{
let mut flows = self.pending_oidc_flows.lock().unwrap();
// Cleanup expired entries
let now = Instant::now();
flows.retain(|_, f| now.duration_since(f.created_at).as_secs() < OIDC_FLOW_TTL_SECS);
flows.insert(
state_token.clone(),
PendingOidcFlow {
created_at: now,
pkce_verifier,
nonce: nonce.clone(),
},
);
}
// Store pending flow (auto-expires after 10 min via moka TTL)
self.pending_oidc_flows.insert(
state_token.clone(),
PendingOidcFlow {
pkce_verifier,
nonce: nonce.clone(),
},
);
// Build authorization URL with state, nonce, and PKCE challenge
let authorize_url = oidc
@@ -952,28 +942,15 @@ impl AuthApplicationService {
/// issue internal tokens, and return a one-time exchange code.
pub async fn oidc_callback(&self, code: &str, state: &str) -> Result<String, DomainError> {
// 0. Validate CSRF state and retrieve PKCE verifier + nonce
let (pkce_verifier, nonce) = {
let mut flows = self.pending_oidc_flows.lock().unwrap();
let flow = flows.remove(state).ok_or_else(|| {
tracing::warn!("OIDC callback with invalid/expired state token");
DomainError::new(
ErrorKind::AccessDenied, "OIDC",
"Invalid or expired OIDC state — possible CSRF attack. Please try logging in again.",
)
})?;
// Check TTL
if Instant::now().duration_since(flow.created_at).as_secs() >= OIDC_FLOW_TTL_SECS {
tracing::warn!("OIDC callback with expired state token");
return Err(DomainError::new(
ErrorKind::AccessDenied,
"OIDC",
"OIDC authorization flow expired. Please try logging in again.",
));
}
(flow.pkce_verifier, flow.nonce)
};
// (entry is auto-expired by moka TTL — remove returns None if expired)
let flow = self.pending_oidc_flows.remove(state).ok_or_else(|| {
tracing::warn!("OIDC callback with invalid/expired state token");
DomainError::new(
ErrorKind::AccessDenied, "OIDC",
"Invalid or expired OIDC state — possible CSRF attack. Please try logging in again.",
)
})?;
let (pkce_verifier, nonce) = (flow.pkce_verifier, flow.nonce);
// Clone the Arc and config out of the RwLock so we don't hold the lock across await points
let (oidc, oidc_config) = {
@@ -1178,20 +1155,11 @@ impl AuthApplicationService {
OsRng.fill_bytes(&mut code_bytes);
let exchange_code = hex::encode(code_bytes);
{
let mut tokens = self.pending_oidc_tokens.lock().unwrap();
// Cleanup expired entries
let now = Instant::now();
tokens.retain(|_, t| now.duration_since(t.created_at).as_secs() < OIDC_TOKEN_TTL_SECS);
tokens.insert(
exchange_code.clone(),
PendingOidcToken {
auth_response,
created_at: now,
},
);
}
// Store auth response (auto-expires after 60 s via moka TTL)
self.pending_oidc_tokens.insert(
exchange_code.clone(),
PendingOidcToken { auth_response },
);
tracing::info!("OIDC login successful, one-time exchange code generated");
@@ -1199,10 +1167,9 @@ impl AuthApplicationService {
}
/// Exchange a one-time code for the authentication tokens.
/// The code is single-use and expires after 60 seconds.
/// The code is single-use and expires after 60 seconds (moka TTL).
pub fn exchange_oidc_token(&self, one_time_code: &str) -> Result<AuthResponseDto, DomainError> {
let mut tokens = self.pending_oidc_tokens.lock().unwrap();
let pending = tokens.remove(one_time_code).ok_or_else(|| {
let pending = self.pending_oidc_tokens.remove(one_time_code).ok_or_else(|| {
DomainError::new(
ErrorKind::AccessDenied,
"OIDC",
@@ -1210,15 +1177,6 @@ impl AuthApplicationService {
)
})?;
// Check TTL
if Instant::now().duration_since(pending.created_at).as_secs() >= OIDC_TOKEN_TTL_SECS {
return Err(DomainError::new(
ErrorKind::AccessDenied,
"OIDC",
"Exchange code expired. Please try logging in again.",
));
}
Ok(pending.auth_response)
}
+29 -9
View File
@@ -1,4 +1,4 @@
use futures::{Future, future::join_all};
use futures::{Future, StreamExt, future::join_all};
use std::sync::Arc;
use thiserror::Error;
use tokio::sync::Semaphore;
@@ -702,18 +702,30 @@ impl BatchOperationService {
// Add individual files at the root of the ZIP
for file_id in &file_ids {
match self.file_retrieval.get_file(file_id).await {
Ok(file_dto) => match self.file_retrieval.get_file_content(file_id).await {
Ok(content) => {
Ok(file_dto) => match self.file_retrieval.get_file_stream(file_id).await {
Ok(stream) => {
let mut stream = std::pin::Pin::from(stream);
if let Err(e) = zip.start_file(&file_dto.name, options) {
info!("Could not start zip entry for {}: {}", file_dto.name, e);
continue;
}
if let Err(e) = zip.write_all(&content) {
info!("Could not write zip entry for {}: {}", file_dto.name, e);
while let Some(chunk) = stream.next().await {
match chunk {
Ok(bytes) => {
if let Err(e) = zip.write_all(&bytes) {
info!("Could not write zip chunk for {}: {}", file_dto.name, e);
break;
}
}
Err(e) => {
info!("Stream error for {}: {}", file_dto.name, e);
break;
}
}
}
}
Err(e) => {
info!("Could not read file content {}: {}", file_id, e);
info!("Could not stream file content {}: {}", file_id, e);
}
},
Err(e) => {
@@ -805,13 +817,21 @@ impl BatchOperationService {
let dir_path = format!("{}/", current.path);
let _ = zip.add_directory(&dir_path, *options);
// Add files
// Add files via streaming (constant ~64 KB memory per file)
if let Ok(files) = self.file_retrieval.list_files(Some(&current.id)).await {
for file in files {
let file_path = format!("{}{}", dir_path, file.name);
if let Ok(content) = self.file_retrieval.get_file_content(&file.id).await {
if let Ok(stream) = self.file_retrieval.get_file_stream(&file.id).await {
let mut stream = std::pin::Pin::from(stream);
if zip.start_file(&file_path, *options).is_ok() {
let _ = zip.write_all(&content);
while let Some(chunk) = stream.next().await {
match chunk {
Ok(bytes) => {
let _ = zip.write_all(&bytes);
}
Err(_) => break,
}
}
}
}
}
@@ -1,6 +1,6 @@
use async_trait::async_trait;
use bytes::Bytes;
use futures::Stream;
use bytes::{Bytes, BytesMut};
use futures::{Stream, StreamExt};
use std::sync::Arc;
use crate::application::dtos::file_dto::FileDto;
@@ -9,7 +9,7 @@ use crate::application::ports::file_ports::{FileRetrievalUseCase, OptimizedFileC
use crate::application::ports::storage_ports::FileReadPort;
use crate::application::ports::transcode_ports::{ImageTranscodePort, OutputFormat};
use crate::common::errors::DomainError;
use tracing::{debug, info, warn};
use tracing::{debug, info};
/// Threshold below which files are served from RAM cache (10 MB).
const CACHE_THRESHOLD: u64 = 10 * 1024 * 1024;
@@ -183,10 +183,17 @@ impl FileRetrievalService {
));
}
// Cache miss – load from disk
// Cache miss – load from disk via streaming (constant 64 KB memory)
debug!("💾 TIER 1 Cache MISS: {} – loading from disk", file_name);
let content = self.file_read.get_file_content(id).await?;
let content_bytes = Bytes::from(content);
let stream = self.file_read.get_file_stream(id).await?;
let mut stream = std::pin::Pin::from(stream);
let mut buf = BytesMut::with_capacity(file_size as usize);
while let Some(chunk) = stream.next().await {
buf.extend_from_slice(&chunk.map_err(|e| {
DomainError::internal_error("File", format!("Stream read error: {}", e))
})?);
}
let content_bytes = buf.freeze();
// Store in cache
if let Some(cache) = &self.content_cache {
@@ -231,21 +238,8 @@ impl FileRetrievalService {
file_name,
file_size / (1024 * 1024)
);
match self.file_read.get_file_stream(id).await {
Ok(stream) => Ok((dto, OptimizedFileContent::Stream(Box::into_pin(stream)))),
Err(e) => {
warn!("Streaming failed, last-resort content load: {}", e);
let content = self.file_read.get_file_content(id).await?;
Ok((
dto,
OptimizedFileContent::Bytes {
data: Bytes::from(content),
mime_type: mime_type.clone(),
was_transcoded: false,
},
))
}
}
let stream = self.file_read.get_file_stream(id).await?;
Ok((dto, OptimizedFileContent::Stream(Box::into_pin(stream))))
}
}
@@ -273,10 +267,6 @@ impl FileRetrievalUseCase for FileRetrievalService {
Ok(files.into_iter().map(FileDto::from).collect())
}
async fn get_file_content(&self, id: &str) -> Result<Vec<u8>, DomainError> {
self.file_read.get_file_content(id).await
}
async fn get_file_stream(
&self,
id: &str,
+4 -11
View File
@@ -348,7 +348,7 @@ impl ShareUseCase for ShareService {
// Verify the password using the infrastructure port
match share.password_hash() {
Some(hash) => self.password_hasher.verify_password(password, hash),
Some(hash) => self.password_hasher.verify_password(password, hash).await,
None => Ok(true), // No password required
}
}
@@ -395,12 +395,13 @@ mod tests {
struct MockPasswordHasher;
#[async_trait]
impl PasswordHasherPort for MockPasswordHasher {
fn hash_password(&self, password: &str) -> Result<String, DomainError> {
async fn hash_password(&self, password: &str) -> Result<String, DomainError> {
Ok(format!("hashed_{}", password))
}
fn verify_password(&self, _password: &str, _hash: &str) -> Result<bool, DomainError> {
async fn verify_password(&self, _password: &str, _hash: &str) -> Result<bool, DomainError> {
Ok(true)
}
}
@@ -439,10 +440,6 @@ mod tests {
unimplemented!()
}
async fn get_file_content(&self, _id: &str) -> Result<Vec<u8>, DomainError> {
unimplemented!()
}
async fn get_file_stream(
&self,
_id: &str,
@@ -465,10 +462,6 @@ mod tests {
unimplemented!()
}
async fn get_file_mmap(&self, _id: &str) -> Result<bytes::Bytes, DomainError> {
unimplemented!()
}
async fn get_file_path(
&self,
_id: &str,
@@ -142,10 +142,6 @@ impl FileReadPort for MockFileRepository {
Ok(vec![])
}
async fn get_file_content(&self, _id: &str) -> std::result::Result<Vec<u8>, DomainError> {
Ok(vec![])
}
async fn get_file_stream(
&self,
_id: &str,
@@ -168,10 +164,6 @@ impl FileReadPort for MockFileRepository {
unimplemented!()
}
async fn get_file_mmap(&self, _id: &str) -> std::result::Result<Bytes, DomainError> {
unimplemented!()
}
async fn get_file_path(&self, _id: &str) -> std::result::Result<StoragePath, DomainError> {
unimplemented!()
}