style: cargo fmt --all

This commit is contained in:
Dionisio
2026-03-03 01:49:18 +01:00
parent 1df52fd702
commit efcf88c4d7
29 changed files with 2754 additions and 2732 deletions
+15 -29
View File
@@ -351,27 +351,23 @@ impl CalDavAdapter {
username username
))))?; ))))?;
xml_writer.write_event(Event::End(BytesEnd::new("D:href")))?; xml_writer.write_event(Event::End(BytesEnd::new("D:href")))?;
xml_writer xml_writer.write_event(Event::End(BytesEnd::new("D:current-user-principal")))?;
.write_event(Event::End(BytesEnd::new("D:current-user-principal")))?;
// calendar-home-set // calendar-home-set
xml_writer xml_writer.write_event(Event::Start(BytesStart::new("C:calendar-home-set")))?;
.write_event(Event::Start(BytesStart::new("C:calendar-home-set")))?;
xml_writer.write_event(Event::Start(BytesStart::new("D:href")))?; xml_writer.write_event(Event::Start(BytesStart::new("D:href")))?;
xml_writer.write_event(Event::Text(BytesText::new(&format!( xml_writer.write_event(Event::Text(BytesText::new(&format!(
"/caldav/{}/", "/caldav/{}/",
username username
))))?; ))))?;
xml_writer.write_event(Event::End(BytesEnd::new("D:href")))?; xml_writer.write_event(Event::End(BytesEnd::new("D:href")))?;
xml_writer xml_writer.write_event(Event::End(BytesEnd::new("C:calendar-home-set")))?;
.write_event(Event::End(BytesEnd::new("C:calendar-home-set")))?;
} }
PropFindType::PropName => { PropFindType::PropName => {
xml_writer.write_event(Event::Empty(BytesStart::new("D:resourcetype")))?; xml_writer.write_event(Event::Empty(BytesStart::new("D:resourcetype")))?;
xml_writer xml_writer
.write_event(Event::Empty(BytesStart::new("D:current-user-principal")))?; .write_event(Event::Empty(BytesStart::new("D:current-user-principal")))?;
xml_writer xml_writer.write_event(Event::Empty(BytesStart::new("C:calendar-home-set")))?;
.write_event(Event::Empty(BytesStart::new("C:calendar-home-set")))?;
} }
PropFindType::Prop(props) => { PropFindType::Prop(props) => {
Self::write_root_requested_props(xml_writer, username, props)?; Self::write_root_requested_props(xml_writer, username, props)?;
@@ -404,9 +400,8 @@ impl CalDavAdapter {
xml_writer.write_event(Event::End(BytesEnd::new("D:resourcetype")))?; xml_writer.write_event(Event::End(BytesEnd::new("D:resourcetype")))?;
} }
("DAV:", "current-user-principal") => { ("DAV:", "current-user-principal") => {
xml_writer.write_event(Event::Start(BytesStart::new( xml_writer
"D:current-user-principal", .write_event(Event::Start(BytesStart::new("D:current-user-principal")))?;
)))?;
xml_writer.write_event(Event::Start(BytesStart::new("D:href")))?; xml_writer.write_event(Event::Start(BytesStart::new("D:href")))?;
xml_writer.write_event(Event::Text(BytesText::new(&format!( xml_writer.write_event(Event::Text(BytesText::new(&format!(
"/caldav/principals/{}/", "/caldav/principals/{}/",
@@ -417,16 +412,14 @@ impl CalDavAdapter {
.write_event(Event::End(BytesEnd::new("D:current-user-principal")))?; .write_event(Event::End(BytesEnd::new("D:current-user-principal")))?;
} }
("urn:ietf:params:xml:ns:caldav", "calendar-home-set") => { ("urn:ietf:params:xml:ns:caldav", "calendar-home-set") => {
xml_writer xml_writer.write_event(Event::Start(BytesStart::new("C:calendar-home-set")))?;
.write_event(Event::Start(BytesStart::new("C:calendar-home-set")))?;
xml_writer.write_event(Event::Start(BytesStart::new("D:href")))?; xml_writer.write_event(Event::Start(BytesStart::new("D:href")))?;
xml_writer.write_event(Event::Text(BytesText::new(&format!( xml_writer.write_event(Event::Text(BytesText::new(&format!(
"/caldav/{}/", "/caldav/{}/",
username username
))))?; ))))?;
xml_writer.write_event(Event::End(BytesEnd::new("D:href")))?; xml_writer.write_event(Event::End(BytesEnd::new("D:href")))?;
xml_writer xml_writer.write_event(Event::End(BytesEnd::new("C:calendar-home-set")))?;
.write_event(Event::End(BytesEnd::new("C:calendar-home-set")))?;
} }
("DAV:", "displayname") => { ("DAV:", "displayname") => {
xml_writer.write_event(Event::Start(BytesStart::new("D:displayname")))?; xml_writer.write_event(Event::Start(BytesStart::new("D:displayname")))?;
@@ -452,10 +445,7 @@ impl CalDavAdapter {
} }
/// Write standard properties for a principal resource. /// Write standard properties for a principal resource.
fn write_principal_props<W: Write>( fn write_principal_props<W: Write>(xml_writer: &mut Writer<W>, username: &str) -> Result<()> {
xml_writer: &mut Writer<W>,
username: &str,
) -> Result<()> {
// resourcetype — principal // resourcetype — principal
xml_writer.write_event(Event::Start(BytesStart::new("D:resourcetype")))?; xml_writer.write_event(Event::Start(BytesStart::new("D:resourcetype")))?;
xml_writer.write_event(Event::Empty(BytesStart::new("D:collection")))?; xml_writer.write_event(Event::Empty(BytesStart::new("D:collection")))?;
@@ -510,9 +500,8 @@ impl CalDavAdapter {
xml_writer.write_event(Event::End(BytesEnd::new("D:displayname")))?; xml_writer.write_event(Event::End(BytesEnd::new("D:displayname")))?;
} }
("DAV:", "current-user-principal") => { ("DAV:", "current-user-principal") => {
xml_writer.write_event(Event::Start(BytesStart::new( xml_writer
"D:current-user-principal", .write_event(Event::Start(BytesStart::new("D:current-user-principal")))?;
)))?;
xml_writer.write_event(Event::Start(BytesStart::new("D:href")))?; xml_writer.write_event(Event::Start(BytesStart::new("D:href")))?;
xml_writer.write_event(Event::Text(BytesText::new(&format!( xml_writer.write_event(Event::Text(BytesText::new(&format!(
"/caldav/principals/{}/", "/caldav/principals/{}/",
@@ -523,16 +512,14 @@ impl CalDavAdapter {
.write_event(Event::End(BytesEnd::new("D:current-user-principal")))?; .write_event(Event::End(BytesEnd::new("D:current-user-principal")))?;
} }
("urn:ietf:params:xml:ns:caldav", "calendar-home-set") => { ("urn:ietf:params:xml:ns:caldav", "calendar-home-set") => {
xml_writer xml_writer.write_event(Event::Start(BytesStart::new("C:calendar-home-set")))?;
.write_event(Event::Start(BytesStart::new("C:calendar-home-set")))?;
xml_writer.write_event(Event::Start(BytesStart::new("D:href")))?; xml_writer.write_event(Event::Start(BytesStart::new("D:href")))?;
xml_writer.write_event(Event::Text(BytesText::new(&format!( xml_writer.write_event(Event::Text(BytesText::new(&format!(
"/caldav/{}/", "/caldav/{}/",
username username
))))?; ))))?;
xml_writer.write_event(Event::End(BytesEnd::new("D:href")))?; xml_writer.write_event(Event::End(BytesEnd::new("D:href")))?;
xml_writer xml_writer.write_event(Event::End(BytesEnd::new("C:calendar-home-set")))?;
.write_event(Event::End(BytesEnd::new("C:calendar-home-set")))?;
} }
("urn:ietf:params:xml:ns:caldav", "calendar-user-address-set") => { ("urn:ietf:params:xml:ns:caldav", "calendar-user-address-set") => {
xml_writer.write_event(Event::Start(BytesStart::new( xml_writer.write_event(Event::Start(BytesStart::new(
@@ -544,9 +531,8 @@ impl CalDavAdapter {
username username
))))?; ))))?;
xml_writer.write_event(Event::End(BytesEnd::new("D:href")))?; xml_writer.write_event(Event::End(BytesEnd::new("D:href")))?;
xml_writer.write_event(Event::End(BytesEnd::new( xml_writer
"C:calendar-user-address-set", .write_event(Event::End(BytesEnd::new("C:calendar-user-address-set")))?;
)))?;
} }
_ => { _ => {
let prop_name = if prop.namespace == "http://calendarserver.org/ns/" { let prop_name = if prop.namespace == "http://calendarserver.org/ns/" {
+13 -11
View File
@@ -356,7 +356,11 @@ mod tests {
</D:propfind>"#; </D:propfind>"#;
let result = WebDavAdapter::parse_propfind(Cursor::new(xml)); let result = WebDavAdapter::parse_propfind(Cursor::new(xml));
assert!(result.is_ok(), "Failed to parse PROPFIND: {:?}", result.err()); assert!(
result.is_ok(),
"Failed to parse PROPFIND: {:?}",
result.err()
);
let request = result.unwrap(); let request = result.unwrap();
match request.prop_find_type { match request.prop_find_type {
@@ -431,7 +435,11 @@ mod tests {
"/caldav/", "/caldav/",
"testuser", "testuser",
); );
assert!(result.is_ok(), "Failed to generate root propfind: {:?}", result.err()); assert!(
result.is_ok(),
"Failed to generate root propfind: {:?}",
result.err()
);
let xml_str = String::from_utf8(output).expect("Invalid UTF-8"); let xml_str = String::from_utf8(output).expect("Invalid UTF-8");
@@ -511,11 +519,8 @@ mod tests {
}; };
let mut output = Vec::new(); let mut output = Vec::new();
let result = CalDavAdapter::generate_principal_propfind_response( let result =
&mut output, CalDavAdapter::generate_principal_propfind_response(&mut output, &request, "testuser");
&request,
"testuser",
);
assert!(result.is_ok(), "Failed: {:?}", result.err()); assert!(result.is_ok(), "Failed: {:?}", result.err());
let xml_str = String::from_utf8(output).expect("Invalid UTF-8"); let xml_str = String::from_utf8(output).expect("Invalid UTF-8");
@@ -576,10 +581,7 @@ mod tests {
); );
// displayname should be populated // displayname should be populated
assert!( assert!(xml_str.contains("Personal"), "Should contain calendar name");
xml_str.contains("Personal"),
"Should contain calendar name"
);
// supported-calendar-component-set should have VEVENT // supported-calendar-component-set should have VEVENT
assert!( assert!(
+334 -339
View File
@@ -1,339 +1,334 @@
//! App Password application service. //! App Password application service.
//! //!
//! Orchestrates creation, verification, listing, and revocation of //! Orchestrates creation, verification, listing, and revocation of
//! application-specific passwords for DAV clients. //! application-specific passwords for DAV clients.
use crate::application::dtos::app_password_dto::*; use crate::application::dtos::app_password_dto::*;
use crate::application::ports::auth_ports::{ use crate::application::ports::auth_ports::{
AppPasswordStoragePort, PasswordHasherPort, UserStoragePort, AppPasswordStoragePort, PasswordHasherPort, UserStoragePort,
}; };
use crate::common::errors::DomainError; use crate::common::errors::DomainError;
use crate::domain::entities::app_password::AppPassword; use crate::domain::entities::app_password::AppPassword;
use chrono::{Duration, Utc}; use chrono::{Duration, Utc};
use moka::future::Cache; use moka::future::Cache;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration as StdDuration; use std::time::Duration as StdDuration;
/// App password token length (32 random alphanumeric chars after prefix). /// App password token length (32 random alphanumeric chars after prefix).
const TOKEN_LENGTH: usize = 32; const TOKEN_LENGTH: usize = 32;
/// Prefix for all app password tokens (makes them easily identifiable). /// Prefix for all app password tokens (makes them easily identifiable).
const TOKEN_PREFIX: &str = "oxicloud-"; const TOKEN_PREFIX: &str = "oxicloud-";
/// TTL for cached Basic Auth verification results. /// TTL for cached Basic Auth verification results.
/// Balances performance (avoids repeated Argon2id + DB queries) with security /// Balances performance (avoids repeated Argon2id + DB queries) with security
/// (limits the window during which a revoked app password remains usable). /// (limits the window during which a revoked app password remains usable).
const BASIC_AUTH_CACHE_TTL_SECS: u64 = 30; const BASIC_AUTH_CACHE_TTL_SECS: u64 = 30;
/// Maximum number of cached Basic Auth verifications. /// Maximum number of cached Basic Auth verifications.
/// Each entry is ~160 bytes (32-byte key + 4 small strings), so 10 000 /// Each entry is ~160 bytes (32-byte key + 4 small strings), so 10 000
/// entries ≈ 1.6 MB — negligible compared to other in-memory caches. /// entries ≈ 1.6 MB — negligible compared to other in-memory caches.
const BASIC_AUTH_CACHE_MAX_ENTRIES: u64 = 10_000; const BASIC_AUTH_CACHE_MAX_ENTRIES: u64 = 10_000;
/// Cached identity returned after a successful Basic Auth verification. /// Cached identity returned after a successful Basic Auth verification.
#[derive(Clone)] #[derive(Clone)]
struct CachedBasicAuthResult { struct CachedBasicAuthResult {
user_id: String, user_id: String,
username: String, username: String,
email: String, email: String,
role: String, role: String,
} }
pub struct AppPasswordService { pub struct AppPasswordService {
repo: Arc<dyn AppPasswordStoragePort>, repo: Arc<dyn AppPasswordStoragePort>,
hasher: Arc<dyn PasswordHasherPort>, hasher: Arc<dyn PasswordHasherPort>,
user_repo: Arc<dyn UserStoragePort>, user_repo: Arc<dyn UserStoragePort>,
base_url: String, base_url: String,
/// In-memory cache of successful Basic Auth verifications. /// In-memory cache of successful Basic Auth verifications.
/// ///
/// **Key**: `blake3(username + ":" + password)` — the plain-text password /// **Key**: `blake3(username + ":" + password)` — the plain-text password
/// is never stored; only a cryptographic hash is kept as lookup key. /// is never stored; only a cryptographic hash is kept as lookup key.
/// ///
/// **Value**: the authenticated identity (user_id, username, email, role). /// **Value**: the authenticated identity (user_id, username, email, role).
/// ///
/// **Eviction**: TTL-based (30 s) + capacity-based (10 000 entries). /// **Eviction**: TTL-based (30 s) + capacity-based (10 000 entries).
/// Failed verifications are *never* cached, so brute-force attackers /// Failed verifications are *never* cached, so brute-force attackers
/// always pay the full Argon2id cost. /// always pay the full Argon2id cost.
auth_cache: Cache<[u8; 32], CachedBasicAuthResult>, auth_cache: Cache<[u8; 32], CachedBasicAuthResult>,
} }
impl AppPasswordService { impl AppPasswordService {
pub fn new( pub fn new(
repo: Arc<dyn AppPasswordStoragePort>, repo: Arc<dyn AppPasswordStoragePort>,
hasher: Arc<dyn PasswordHasherPort>, hasher: Arc<dyn PasswordHasherPort>,
user_repo: Arc<dyn UserStoragePort>, user_repo: Arc<dyn UserStoragePort>,
base_url: String, base_url: String,
) -> Self { ) -> Self {
let auth_cache = Cache::builder() let auth_cache = Cache::builder()
.max_capacity(BASIC_AUTH_CACHE_MAX_ENTRIES) .max_capacity(BASIC_AUTH_CACHE_MAX_ENTRIES)
.time_to_live(StdDuration::from_secs(BASIC_AUTH_CACHE_TTL_SECS)) .time_to_live(StdDuration::from_secs(BASIC_AUTH_CACHE_TTL_SECS))
.build(); .build();
tracing::info!( tracing::info!(
"AppPasswordService Basic Auth cache initialized: TTL={}s, max={} entries", "AppPasswordService Basic Auth cache initialized: TTL={}s, max={} entries",
BASIC_AUTH_CACHE_TTL_SECS, BASIC_AUTH_CACHE_TTL_SECS,
BASIC_AUTH_CACHE_MAX_ENTRIES, BASIC_AUTH_CACHE_MAX_ENTRIES,
); );
Self { Self {
repo, repo,
hasher, hasher,
user_repo, user_repo,
base_url, base_url,
auth_cache, auth_cache,
} }
} }
/// Generate a random app password token using cryptographic RNG. /// Generate a random app password token using cryptographic RNG.
fn generate_token() -> String { fn generate_token() -> String {
use rand_core::{OsRng, RngCore}; use rand_core::{OsRng, RngCore};
let charset: &[u8] = b"abcdefghijklmnopqrstuvwxyz\ let charset: &[u8] = b"abcdefghijklmnopqrstuvwxyz\
ABCDEFGHIJKLMNOPQRSTUVWXYZ\ ABCDEFGHIJKLMNOPQRSTUVWXYZ\
0123456789"; 0123456789";
let mut rng_bytes = [0u8; TOKEN_LENGTH]; let mut rng_bytes = [0u8; TOKEN_LENGTH];
OsRng.fill_bytes(&mut rng_bytes); OsRng.fill_bytes(&mut rng_bytes);
let random_part: String = rng_bytes let random_part: String = rng_bytes
.iter() .iter()
.map(|&b| { .map(|&b| {
let idx = (b as usize) % charset.len(); let idx = (b as usize) % charset.len();
charset[idx] as char charset[idx] as char
}) })
.collect(); .collect();
format!("{}{}", TOKEN_PREFIX, random_part) format!("{}{}", TOKEN_PREFIX, random_part)
} }
/// Create a new app password for the given user. /// Create a new app password for the given user.
/// ///
/// Returns the response DTO that includes the plain-text password (shown only once). /// Returns the response DTO that includes the plain-text password (shown only once).
pub async fn create( pub async fn create(
&self, &self,
user_id: &str, user_id: &str,
request: CreateAppPasswordRequestDto, request: CreateAppPasswordRequestDto,
) -> Result<AppPasswordCreatedResponseDto, DomainError> { ) -> Result<AppPasswordCreatedResponseDto, DomainError> {
// Validate label // Validate label
let label = request.label.trim().to_string(); let label = request.label.trim().to_string();
if label.is_empty() || label.len() > 255 { if label.is_empty() || label.len() > 255 {
return Err(DomainError::validation_error( return Err(DomainError::validation_error(
"Label must be 1-255 characters", "Label must be 1-255 characters",
)); ));
} }
// Fetch user for the username (needed for Basic Auth instructions) // Fetch user for the username (needed for Basic Auth instructions)
let user = self.user_repo.get_user_by_id(user_id).await?; let user = self.user_repo.get_user_by_id(user_id).await?;
let username = user.username().to_string(); let username = user.username().to_string();
// Generate the plain-text token // Generate the plain-text token
let plain_token = Self::generate_token(); let plain_token = Self::generate_token();
let prefix = plain_token[..TOKEN_PREFIX.len() + 8].to_string(); let prefix = plain_token[..TOKEN_PREFIX.len() + 8].to_string();
// Hash the token for storage // Hash the token for storage
let password_hash = self.hasher.hash_password(&plain_token).await?; let password_hash = self.hasher.hash_password(&plain_token).await?;
// Calculate expiration // Calculate expiration
let expires_at = request.expires_in_days.map(|days| { let expires_at = request
Utc::now() + Duration::days(days as i64) .expires_in_days
}); .map(|days| Utc::now() + Duration::days(days as i64));
// Create entity // Create entity
let app_password = AppPassword::new( let app_password = AppPassword::new(
user_id.to_string(), user_id.to_string(),
label.clone(), label.clone(),
password_hash, password_hash,
prefix.clone(), prefix.clone(),
request.scopes.clone(), request.scopes.clone(),
expires_at, expires_at,
); );
let saved = self.repo.create(app_password).await?; let saved = self.repo.create(app_password).await?;
let expires_str = saved let expires_str = saved.expires_at.map(|dt| dt.to_rfc3339());
.expires_at
.map(|dt| dt.to_rfc3339()); let curl_example = format!(
"curl -u '{}:{}' -X PROPFIND {}/webdav/",
let curl_example = format!( username, plain_token, self.base_url
"curl -u '{}:{}' -X PROPFIND {}/webdav/", );
username, plain_token, self.base_url
); Ok(AppPasswordCreatedResponseDto {
id: saved.id,
Ok(AppPasswordCreatedResponseDto { label,
id: saved.id, password: plain_token,
label, username: username.clone(),
password: plain_token, scopes: request.scopes,
username: username.clone(), expires_at: expires_str,
scopes: request.scopes, instructions: AppPasswordInstructions {
expires_at: expires_str, davx5: format!(
instructions: AppPasswordInstructions { "In DAVx⁵, add account with base URL: {}/webdav/\n\
davx5: format!( Username: {}\n\
"In DAVx⁵, add account with base URL: {}/webdav/\n\ Password: (the token shown above)",
Username: {}\n\ self.base_url, username
Password: (the token shown above)", ),
self.base_url, username thunderbird: format!(
), "In Thunderbird CalDAV/CardDAV:\n\
thunderbird: format!( URL: {}/caldav/ or {}/carddav/\n\
"In Thunderbird CalDAV/CardDAV:\n\ Username: {}\n\
URL: {}/caldav/ or {}/carddav/\n\ Password: (the token shown above)",
Username: {}\n\ self.base_url, self.base_url, username
Password: (the token shown above)", ),
self.base_url, self.base_url, username rclone: format!(
), "rclone config:\n\
rclone: format!( type = webdav\n\
"rclone config:\n\ url = {}/webdav/\n\
type = webdav\n\ vendor = other\n\
url = {}/webdav/\n\ user = {}\n\
vendor = other\n\ pass = (the token shown above, use 'rclone obscure' to encode)",
user = {}\n\ self.base_url, username
pass = (the token shown above, use 'rclone obscure' to encode)", ),
self.base_url, username curl_example,
), },
curl_example, })
}, }
})
} /// List all app passwords for a user (excludes plain-text passwords).
pub async fn list(&self, user_id: &str) -> Result<AppPasswordListResponseDto, DomainError> {
/// List all app passwords for a user (excludes plain-text passwords). let passwords = self.repo.list_by_user(user_id).await?;
pub async fn list(&self, user_id: &str) -> Result<AppPasswordListResponseDto, DomainError> { let total = passwords.len();
let passwords = self.repo.list_by_user(user_id).await?;
let total = passwords.len(); let app_passwords = passwords
.into_iter()
let app_passwords = passwords .map(|ap| {
.into_iter() let is_active = ap.active && !ap.is_expired();
.map(|ap| { AppPasswordSummaryDto {
let is_active = ap.active && !ap.is_expired(); id: ap.id,
AppPasswordSummaryDto { label: ap.label,
id: ap.id, prefix: format!("{}...", ap.prefix),
label: ap.label, scopes: ap.scopes,
prefix: format!("{}...", ap.prefix), created_at: ap.created_at.to_rfc3339(),
scopes: ap.scopes, last_used_at: ap.last_used_at.map(|dt| dt.to_rfc3339()),
created_at: ap.created_at.to_rfc3339(), expires_at: ap.expires_at.map(|dt| dt.to_rfc3339()),
last_used_at: ap.last_used_at.map(|dt| dt.to_rfc3339()), active: is_active,
expires_at: ap.expires_at.map(|dt| dt.to_rfc3339()), }
active: is_active, })
} .collect();
})
.collect(); Ok(AppPasswordListResponseDto {
app_passwords,
Ok(AppPasswordListResponseDto { total,
app_passwords, })
total, }
})
} /// Revoke (soft-delete) an app password. Verifies ownership.
///
/// Revoke (soft-delete) an app password. Verifies ownership. /// Also invalidates **all** cached Basic Auth entries for the owning user
/// /// so that the revocation takes effect immediately (instead of waiting
/// Also invalidates **all** cached Basic Auth entries for the owning user /// up to `BASIC_AUTH_CACHE_TTL_SECS`).
/// so that the revocation takes effect immediately (instead of waiting pub async fn revoke(
/// up to `BASIC_AUTH_CACHE_TTL_SECS`). &self,
pub async fn revoke(&self, user_id: &str, id: &str) -> Result<AppPasswordRevokeResponseDto, DomainError> { user_id: &str,
let ap = self.repo.get_by_id(id).await?; id: &str,
if ap.user_id != user_id { ) -> Result<AppPasswordRevokeResponseDto, DomainError> {
return Err(DomainError::unauthorized( let ap = self.repo.get_by_id(id).await?;
"You can only revoke your own app passwords", if ap.user_id != user_id {
)); return Err(DomainError::unauthorized(
} "You can only revoke your own app passwords",
self.repo.revoke(id).await?; ));
}
// Invalidate all cached auth entries for this user so the self.repo.revoke(id).await?;
// revocation is effective immediately.
let uid = user_id.to_string(); // Invalidate all cached auth entries for this user so the
self.auth_cache // revocation is effective immediately.
.invalidate_entries_if(move |_key, val| val.user_id == uid) let uid = user_id.to_string();
.ok(); self.auth_cache
.invalidate_entries_if(move |_key, val| val.user_id == uid)
tracing::debug!("Revoked app password {} — auth cache entries for user {} invalidated", id, user_id); .ok();
Ok(AppPasswordRevokeResponseDto { tracing::debug!(
status: "revoked".to_string(), "Revoked app password {} — auth cache entries for user {} invalidated",
id: id.to_string(), id,
}) user_id
} );
/// Verify username + app password for HTTP Basic Auth. Ok(AppPasswordRevokeResponseDto {
/// status: "revoked".to_string(),
/// Returns `(user_id, username, email, role)` on success. id: id.to_string(),
/// })
/// ## Performance }
///
/// Successful verifications are cached for `BASIC_AUTH_CACHE_TTL_SECS` /// Verify username + app password for HTTP Basic Auth.
/// (default 30 s) keyed by `blake3(username:password)`. This avoids ///
/// the expensive Argon2id computation **and** the three PostgreSQL /// Returns `(user_id, username, email, role)` on success.
/// round-trips on every repeated DAV request from the same client. ///
/// /// ## Performance
/// Failed verifications are **never** cached, preserving the full ///
/// Argon2id cost as a brute-force deterrent. /// Successful verifications are cached for `BASIC_AUTH_CACHE_TTL_SECS`
pub async fn verify_basic_auth( /// (default 30 s) keyed by `blake3(username:password)`. This avoids
&self, /// the expensive Argon2id computation **and** the three PostgreSQL
username: &str, /// round-trips on every repeated DAV request from the same client.
password: &str, ///
) -> Result<(String, String, String, String), DomainError> { /// Failed verifications are **never** cached, preserving the full
// ── 1. Compute cache key = blake3("username:password") ──────── /// Argon2id cost as a brute-force deterrent.
// The plain-text password is never stored; only the 32-byte pub async fn verify_basic_auth(
// cryptographic digest is used as lookup key. &self,
let cache_key: [u8; 32] = blake3::hash( username: &str,
format!("{}:{}", username, password).as_bytes(), password: &str,
) ) -> Result<(String, String, String, String), DomainError> {
.into(); // ── 1. Compute cache key = blake3("username:password") ────────
// The plain-text password is never stored; only the 32-byte
// ── 2. Cache hit → return immediately ──────────────────────── // cryptographic digest is used as lookup key.
if let Some(cached) = self.auth_cache.get(&cache_key).await { let cache_key: [u8; 32] =
return Ok(( blake3::hash(format!("{}:{}", username, password).as_bytes()).into();
cached.user_id,
cached.username, // ── 2. Cache hit → return immediately ────────────────────────
cached.email, if let Some(cached) = self.auth_cache.get(&cache_key).await {
cached.role, return Ok((cached.user_id, cached.username, cached.email, cached.role));
)); }
}
// ── 3. Cache miss → full verification ────────────────────────
// ── 3. Cache miss → full verification ──────────────────────── // Look up user by username
// Look up user by username let user = self
let user = self .user_repo
.user_repo .get_user_by_username(username)
.get_user_by_username(username) .await
.await .map_err(|_| DomainError::unauthorized("Invalid username or app password"))?;
.map_err(|_| DomainError::unauthorized("Invalid username or app password"))?;
// Get all active app passwords for this user
// Get all active app passwords for this user let app_passwords = self.repo.get_active_by_user_id(user.id()).await?;
let app_passwords = self
.repo if app_passwords.is_empty() {
.get_active_by_user_id(user.id()) return Err(DomainError::unauthorized(
.await?; "Invalid username or app password",
));
if app_passwords.is_empty() { }
return Err(DomainError::unauthorized(
"Invalid username or app password", // Try each app password hash (Argon2id — CPU-intensive)
)); for ap in &app_passwords {
} if let Ok(true) = self
.hasher
// Try each app password hash (Argon2id — CPU-intensive) .verify_password(password, &ap.password_hash)
for ap in &app_passwords { .await
if let Ok(true) = self.hasher.verify_password(password, &ap.password_hash).await { {
// Update last_used_at (fire-and-forget; don't fail auth on touch error) // Update last_used_at (fire-and-forget; don't fail auth on touch error)
let _ = self.repo.touch_last_used(&ap.id).await; let _ = self.repo.touch_last_used(&ap.id).await;
let result = CachedBasicAuthResult { let result = CachedBasicAuthResult {
user_id: user.id().to_string(), user_id: user.id().to_string(),
username: user.username().to_string(), username: user.username().to_string(),
email: user.email().to_string(), email: user.email().to_string(),
role: user.role().to_string(), role: user.role().to_string(),
}; };
// ── 4. Cache the successful result ──────────────────── // ── 4. Cache the successful result ────────────────────
self.auth_cache.insert(cache_key, result.clone()).await; self.auth_cache.insert(cache_key, result.clone()).await;
return Ok(( return Ok((result.user_id, result.username, result.email, result.role));
result.user_id, }
result.username, }
result.email,
result.role, // Failed verifications are intentionally NOT cached so that
)); // brute-force attackers always pay the full Argon2id cost.
} Err(DomainError::unauthorized(
} "Invalid username or app password",
))
// Failed verifications are intentionally NOT cached so that }
// brute-force attackers always pay the full Argon2id cost. }
Err(DomainError::unauthorized(
"Invalid username or app password",
))
}
}
+9 -3
View File
@@ -138,7 +138,9 @@ impl BatchOperationService {
let target_folder = target_folder.clone(); let target_folder = target_folder.clone();
async move { async move {
let copy_result = mgmt.copy_file(&file_id, target_folder.map(|s| s.to_string())).await; let copy_result = mgmt
.copy_file(&file_id, target_folder.map(|s| s.to_string()))
.await;
(file_id, copy_result) (file_id, copy_result)
} }
})) }))
@@ -201,7 +203,9 @@ impl BatchOperationService {
let target_folder = target_folder.clone(); let target_folder = target_folder.clone();
async move { async move {
let move_result = mgmt.move_file(&file_id, target_folder.map(|s| s.to_string())).await; let move_result = mgmt
.move_file(&file_id, target_folder.map(|s| s.to_string()))
.await;
(file_id, move_result) (file_id, move_result)
} }
})) }))
@@ -578,7 +582,9 @@ impl BatchOperationService {
let caller = caller.clone(); let caller = caller.clone();
async move { async move {
let dto = MoveFolderDto { parent_id: target.map(|s| s.to_string()) }; let dto = MoveFolderDto {
parent_id: target.map(|s| s.to_string()),
};
let move_result = folder_service.move_folder(&folder_id, dto, &caller).await; let move_result = folder_service.move_folder(&folder_id, dto, &caller).await;
(folder_id, move_result) (folder_id, move_result)
} }
+435 -437
View File
@@ -1,437 +1,435 @@
//! OAuth 2.0 Device Authorization Grant service (RFC 8628). //! OAuth 2.0 Device Authorization Grant service (RFC 8628).
//! //!
//! Orchestrates the full device flow: //! Orchestrates the full device flow:
//! 1. `initiate` — generates device_code + user_code, stores in DB //! 1. `initiate` — generates device_code + user_code, stores in DB
//! 2. `verify_user_code` — looks up pending code for the verification page //! 2. `verify_user_code` — looks up pending code for the verification page
//! 3. `approve` — user approves, tokens are generated and stored //! 3. `approve` — user approves, tokens are generated and stored
//! 4. `deny` — user denies the request //! 4. `deny` — user denies the request
//! 5. `poll` — client polls by device_code; returns tokens or status error //! 5. `poll` — client polls by device_code; returns tokens or status error
//! 6. `cleanup_expired` — background job to purge stale entries //! 6. `cleanup_expired` — background job to purge stale entries
use std::sync::Arc; use std::sync::Arc;
use crate::application::dtos::device_auth_dto::*; use crate::application::dtos::device_auth_dto::*;
use crate::application::ports::auth_ports::{DeviceCodeStoragePort, TokenServicePort, UserStoragePort}; use crate::application::ports::auth_ports::SessionStoragePort;
use crate::common::errors::{DomainError, ErrorKind}; use crate::application::ports::auth_ports::{
use crate::domain::entities::device_code::{DeviceCode, DeviceCodeStatus}; DeviceCodeStoragePort, TokenServicePort, UserStoragePort,
use crate::domain::entities::session::Session; };
use crate::application::ports::auth_ports::SessionStoragePort; use crate::common::errors::{DomainError, ErrorKind};
use crate::domain::entities::device_code::{DeviceCode, DeviceCodeStatus};
/// Default device code lifetime: 15 minutes (RFC 8628 recommends 5-30 min). use crate::domain::entities::session::Session;
const DEVICE_CODE_LIFETIME_SECS: i64 = 900;
/// Default device code lifetime: 15 minutes (RFC 8628 recommends 5-30 min).
/// Default polling interval in seconds (RFC 8628 §3.2 recommends 5s). const DEVICE_CODE_LIFETIME_SECS: i64 = 900;
const DEFAULT_POLL_INTERVAL: i32 = 5;
/// Default polling interval in seconds (RFC 8628 §3.2 recommends 5s).
/// Length of the device_code (hex-encoded, 64 chars = 32 bytes). const DEFAULT_POLL_INTERVAL: i32 = 5;
const DEVICE_CODE_BYTES: usize = 32;
/// Length of the device_code (hex-encoded, 64 chars = 32 bytes).
/// User code format: 4 uppercase letters + hyphen + 4 digits → "ABCD-1234" const DEVICE_CODE_BYTES: usize = 32;
/// Short enough to type, long enough to avoid collisions with 26^4 * 10^4 = ~4.5 billion combos.
const USER_CODE_LETTER_LEN: usize = 4; /// User code format: 4 uppercase letters + hyphen + 4 digits → "ABCD-1234"
const USER_CODE_DIGIT_LEN: usize = 4; /// Short enough to type, long enough to avoid collisions with 26^4 * 10^4 = ~4.5 billion combos.
const USER_CODE_LETTER_LEN: usize = 4;
pub struct DeviceAuthService { const USER_CODE_DIGIT_LEN: usize = 4;
device_code_storage: Arc<dyn DeviceCodeStoragePort>,
token_service: Arc<dyn TokenServicePort>, pub struct DeviceAuthService {
user_storage: Arc<dyn UserStoragePort>, device_code_storage: Arc<dyn DeviceCodeStoragePort>,
session_storage: Arc<dyn SessionStoragePort>, token_service: Arc<dyn TokenServicePort>,
/// Base URL of the server (e.g. "https://cloud.example.com") user_storage: Arc<dyn UserStoragePort>,
base_url: String, session_storage: Arc<dyn SessionStoragePort>,
} /// Base URL of the server (e.g. "https://cloud.example.com")
base_url: String,
impl DeviceAuthService { }
pub fn new(
device_code_storage: Arc<dyn DeviceCodeStoragePort>, impl DeviceAuthService {
token_service: Arc<dyn TokenServicePort>, pub fn new(
user_storage: Arc<dyn UserStoragePort>, device_code_storage: Arc<dyn DeviceCodeStoragePort>,
session_storage: Arc<dyn SessionStoragePort>, token_service: Arc<dyn TokenServicePort>,
base_url: String, user_storage: Arc<dyn UserStoragePort>,
) -> Self { session_storage: Arc<dyn SessionStoragePort>,
Self { base_url: String,
device_code_storage, ) -> Self {
token_service, Self {
user_storage, device_code_storage,
session_storage, token_service,
base_url, user_storage,
} session_storage,
} base_url,
}
// ======================================================================== }
// 1. Initiate — called by the DAV client
// ======================================================================== // ========================================================================
// 1. Initiate — called by the DAV client
/// Start a new device authorization flow. // ========================================================================
///
/// Returns the response that the client displays to the user. /// Start a new device authorization flow.
pub async fn initiate( ///
&self, /// Returns the response that the client displays to the user.
req: DeviceAuthorizeRequestDto, pub async fn initiate(
) -> Result<DeviceAuthorizeResponseDto, DomainError> { &self,
let device_code_token = generate_device_code(); req: DeviceAuthorizeRequestDto,
let user_code = generate_user_code(); ) -> Result<DeviceAuthorizeResponseDto, DomainError> {
let device_code_token = generate_device_code();
let verification_uri = format!("{}/device", self.base_url.trim_end_matches('/')); let user_code = generate_user_code();
let verification_uri_complete = format!("{}?code={}", verification_uri, user_code);
let verification_uri = format!("{}/device", self.base_url.trim_end_matches('/'));
let dc = DeviceCode::new( let verification_uri_complete = format!("{}?code={}", verification_uri, user_code);
device_code_token.clone(),
user_code.clone(), let dc = DeviceCode::new(
req.client_name, device_code_token.clone(),
req.scope, user_code.clone(),
verification_uri.clone(), req.client_name,
Some(verification_uri_complete.clone()), req.scope,
DEVICE_CODE_LIFETIME_SECS, verification_uri.clone(),
DEFAULT_POLL_INTERVAL, Some(verification_uri_complete.clone()),
); DEVICE_CODE_LIFETIME_SECS,
DEFAULT_POLL_INTERVAL,
let dc = self.device_code_storage.create_device_code(dc).await?; );
tracing::info!( let dc = self.device_code_storage.create_device_code(dc).await?;
"Device auth flow initiated: user_code={}, expires_in={}s",
user_code, tracing::info!(
DEVICE_CODE_LIFETIME_SECS "Device auth flow initiated: user_code={}, expires_in={}s",
); user_code,
DEVICE_CODE_LIFETIME_SECS
Ok(DeviceAuthorizeResponseDto { );
device_code: device_code_token,
user_code, Ok(DeviceAuthorizeResponseDto {
verification_uri, device_code: device_code_token,
verification_uri_complete: Some(verification_uri_complete), user_code,
expires_in: dc.seconds_remaining(), verification_uri,
interval: DEFAULT_POLL_INTERVAL, verification_uri_complete: Some(verification_uri_complete),
}) expires_in: dc.seconds_remaining(),
} interval: DEFAULT_POLL_INTERVAL,
})
// ======================================================================== }
// 2. Verify — user opens the verification page, looks up pending code
// ======================================================================== // ========================================================================
// 2. Verify — user opens the verification page, looks up pending code
/// Look up a pending device code by user_code for the verification page. // ========================================================================
pub async fn verify_user_code(
&self, /// Look up a pending device code by user_code for the verification page.
user_code: &str, pub async fn verify_user_code(
) -> Result<DeviceVerifyInfoDto, DomainError> { &self,
let normalized = user_code.trim().to_uppercase().replace(' ', ""); user_code: &str,
) -> Result<DeviceVerifyInfoDto, DomainError> {
match self let normalized = user_code.trim().to_uppercase().replace(' ', "");
.device_code_storage
.get_pending_by_user_code(&normalized) match self
.await .device_code_storage
{ .get_pending_by_user_code(&normalized)
Ok(dc) => { .await
if dc.is_expired() { {
return Ok(DeviceVerifyInfoDto { Ok(dc) => {
client_name: dc.client_name().to_string(), if dc.is_expired() {
scopes: dc.scopes().to_string(), return Ok(DeviceVerifyInfoDto {
valid: false, client_name: dc.client_name().to_string(),
}); scopes: dc.scopes().to_string(),
} valid: false,
Ok(DeviceVerifyInfoDto { });
client_name: dc.client_name().to_string(), }
scopes: dc.scopes().to_string(), Ok(DeviceVerifyInfoDto {
valid: true, client_name: dc.client_name().to_string(),
}) scopes: dc.scopes().to_string(),
} valid: true,
Err(_) => Ok(DeviceVerifyInfoDto { })
client_name: String::new(), }
scopes: String::new(), Err(_) => Ok(DeviceVerifyInfoDto {
valid: false, client_name: String::new(),
}), scopes: String::new(),
} valid: false,
} }),
}
// ======================================================================== }
// 3. Approve — authenticated user approves the device code
// ======================================================================== // ========================================================================
// 3. Approve — authenticated user approves the device code
/// Approve a device code, generating tokens for the polling client. // ========================================================================
///
/// * `user_code` — the code from the verification page /// Approve a device code, generating tokens for the polling client.
/// * `user_id` — the authenticated user's ID (from session/JWT) ///
pub async fn approve( /// * `user_code` — the code from the verification page
&self, /// * `user_id` — the authenticated user's ID (from session/JWT)
user_code: &str, pub async fn approve(&self, user_code: &str, user_id: &str) -> Result<(), DomainError> {
user_id: &str, let normalized = user_code.trim().to_uppercase().replace(' ', "");
) -> Result<(), DomainError> {
let normalized = user_code.trim().to_uppercase().replace(' ', ""); let mut dc = self
.device_code_storage
let mut dc = self .get_pending_by_user_code(&normalized)
.device_code_storage .await?;
.get_pending_by_user_code(&normalized)
.await?; if dc.is_expired() {
return Err(DomainError::new(
if dc.is_expired() { ErrorKind::AccessDenied,
return Err(DomainError::new( "DeviceCode",
ErrorKind::AccessDenied, "Device code has expired. Please start a new authorization flow.",
"DeviceCode", ));
"Device code has expired. Please start a new authorization flow.", }
));
} // Fetch user to generate tokens
let user = self.user_storage.get_user_by_id(user_id).await?;
// Fetch user to generate tokens
let user = self.user_storage.get_user_by_id(user_id).await?; // Generate internal JWT access token + refresh token
let access_token = self.token_service.generate_access_token(&user)?;
// Generate internal JWT access token + refresh token let refresh_token = self.token_service.generate_refresh_token();
let access_token = self.token_service.generate_access_token(&user)?;
let refresh_token = self.token_service.generate_refresh_token(); // Persist refresh token as a session
let session = Session::new(
// Persist refresh token as a session user_id.to_string(),
let session = Session::new( refresh_token.clone(),
user_id.to_string(), None, // ip_address
refresh_token.clone(), Some(format!("device:{}", dc.client_name())), // user_agent
None, // ip_address self.token_service.refresh_token_expiry_days(),
Some(format!("device:{}", dc.client_name())), // user_agent );
self.token_service.refresh_token_expiry_days(), self.session_storage.create_session(session).await?;
);
self.session_storage.create_session(session).await?; // Store tokens on the device code entity
dc.authorize(user_id.to_string(), access_token, refresh_token);
// Store tokens on the device code entity self.device_code_storage.update_device_code(dc).await?;
dc.authorize(user_id.to_string(), access_token, refresh_token);
self.device_code_storage.update_device_code(dc).await?; tracing::info!(
"Device code approved by user {} (user_code={})",
tracing::info!( user_id,
"Device code approved by user {} (user_code={})", normalized
user_id, );
normalized
); Ok(())
}
Ok(())
} // ========================================================================
// 4. Deny — authenticated user denies the device code
// ======================================================================== // ========================================================================
// 4. Deny — authenticated user denies the device code
// ======================================================================== pub async fn deny(&self, user_code: &str) -> Result<(), DomainError> {
let normalized = user_code.trim().to_uppercase().replace(' ', "");
pub async fn deny(&self, user_code: &str) -> Result<(), DomainError> {
let normalized = user_code.trim().to_uppercase().replace(' ', ""); let mut dc = self
.device_code_storage
let mut dc = self .get_pending_by_user_code(&normalized)
.device_code_storage .await?;
.get_pending_by_user_code(&normalized)
.await?; dc.deny();
self.device_code_storage.update_device_code(dc).await?;
dc.deny();
self.device_code_storage.update_device_code(dc).await?; tracing::info!("Device code denied (user_code={})", normalized);
tracing::info!("Device code denied (user_code={})", normalized); Ok(())
}
Ok(())
} // ========================================================================
// 5. Poll — client polls by device_code for tokens
// ======================================================================== // ========================================================================
// 5. Poll — client polls by device_code for tokens
// ======================================================================== /// Client polls for tokens. Returns:
/// - `Ok(DeviceTokenSuccessDto)` if authorized
/// Client polls for tokens. Returns: /// - `Err` with specific RFC 8628 error codes for pending/slow_down/expired/denied
/// - `Ok(DeviceTokenSuccessDto)` if authorized pub async fn poll(&self, device_code: &str) -> Result<DeviceTokenSuccessDto, DevicePollError> {
/// - `Err` with specific RFC 8628 error codes for pending/slow_down/expired/denied let mut dc = self
pub async fn poll( .device_code_storage
&self, .get_by_device_code(device_code)
device_code: &str, .await
) -> Result<DeviceTokenSuccessDto, DevicePollError> { .map_err(|_| DevicePollError::InvalidDeviceCode)?;
let mut dc = self
.device_code_storage // Check expiry first
.get_by_device_code(device_code) if dc.is_expired() && dc.status() == DeviceCodeStatus::Pending {
.await let mut expired_dc = dc.clone();
.map_err(|_| DevicePollError::InvalidDeviceCode)?; expired_dc.mark_expired();
let _ = self
// Check expiry first .device_code_storage
if dc.is_expired() && dc.status() == DeviceCodeStatus::Pending { .update_device_code(expired_dc)
let mut expired_dc = dc.clone(); .await;
expired_dc.mark_expired(); return Err(DevicePollError::ExpiredToken);
let _ = self.device_code_storage.update_device_code(expired_dc).await; }
return Err(DevicePollError::ExpiredToken);
} match dc.status() {
DeviceCodeStatus::Pending => {
match dc.status() { // Check for slow_down (polling too fast)
DeviceCodeStatus::Pending => { if dc.is_polling_too_fast() {
// Check for slow_down (polling too fast) return Err(DevicePollError::SlowDown);
if dc.is_polling_too_fast() { }
return Err(DevicePollError::SlowDown); // Record this poll
} dc.record_poll();
// Record this poll let _ = self.device_code_storage.update_device_code(dc).await;
dc.record_poll(); Err(DevicePollError::AuthorizationPending)
let _ = self.device_code_storage.update_device_code(dc).await; }
Err(DevicePollError::AuthorizationPending) DeviceCodeStatus::Authorized => {
} let access_token = dc.access_token().unwrap_or_default().to_string();
DeviceCodeStatus::Authorized => { let refresh_token = dc.refresh_token().unwrap_or_default().to_string();
let access_token = dc.access_token().unwrap_or_default().to_string(); let scope = dc.scopes().to_string();
let refresh_token = dc.refresh_token().unwrap_or_default().to_string();
let scope = dc.scopes().to_string(); Ok(DeviceTokenSuccessDto {
access_token,
Ok(DeviceTokenSuccessDto { token_type: "Bearer".to_string(),
access_token, refresh_token,
token_type: "Bearer".to_string(), expires_in: self.token_service.refresh_token_expiry_secs(),
refresh_token, scope,
expires_in: self.token_service.refresh_token_expiry_secs(), })
scope, }
}) DeviceCodeStatus::Denied => Err(DevicePollError::AccessDenied),
} DeviceCodeStatus::Expired => Err(DevicePollError::ExpiredToken),
DeviceCodeStatus::Denied => Err(DevicePollError::AccessDenied), }
DeviceCodeStatus::Expired => Err(DevicePollError::ExpiredToken), }
}
} // ========================================================================
// 6. Cleanup — purge expired entries
// ======================================================================== // ========================================================================
// 6. Cleanup — purge expired entries
// ======================================================================== pub async fn cleanup_expired(&self) -> Result<u64, DomainError> {
let deleted = self.device_code_storage.delete_expired().await?;
pub async fn cleanup_expired(&self) -> Result<u64, DomainError> { if deleted > 0 {
let deleted = self.device_code_storage.delete_expired().await?; tracing::info!("Device code cleanup: {} expired entries removed", deleted);
if deleted > 0 { }
tracing::info!("Device code cleanup: {} expired entries removed", deleted); Ok(deleted)
} }
Ok(deleted)
} // ========================================================================
// 7. List — user's authorized devices (for UI)
// ======================================================================== // ========================================================================
// 7. List — user's authorized devices (for UI)
// ======================================================================== pub async fn list_user_devices(
&self,
pub async fn list_user_devices( user_id: &str,
&self, ) -> Result<Vec<DeviceInfoDto>, DomainError> {
user_id: &str, let codes = self.device_code_storage.list_by_user(user_id).await?;
) -> Result<Vec<DeviceInfoDto>, DomainError> { Ok(codes
let codes = self.device_code_storage.list_by_user(user_id).await?; .into_iter()
Ok(codes .map(|dc| DeviceInfoDto {
.into_iter() id: dc.id().to_string(),
.map(|dc| DeviceInfoDto { client_name: dc.client_name().to_string(),
id: dc.id().to_string(), scopes: dc.scopes().to_string(),
client_name: dc.client_name().to_string(), status: dc.status().as_str().to_string(),
scopes: dc.scopes().to_string(), created_at: dc.created_at().to_rfc3339(),
status: dc.status().as_str().to_string(), authorized_at: dc.authorized_at().map(|t| t.to_rfc3339()),
created_at: dc.created_at().to_rfc3339(), expires_at: dc.expires_at().to_rfc3339(),
authorized_at: dc.authorized_at().map(|t| t.to_rfc3339()), })
expires_at: dc.expires_at().to_rfc3339(), .collect())
}) }
.collect())
} // ========================================================================
// 8. Revoke — user revokes a device authorization
// ======================================================================== // ========================================================================
// 8. Revoke — user revokes a device authorization
// ======================================================================== pub async fn revoke_device(&self, device_id: &str, user_id: &str) -> Result<(), DomainError> {
// Verify ownership before deleting
pub async fn revoke_device(&self, device_id: &str, user_id: &str) -> Result<(), DomainError> { let devices = self.device_code_storage.list_by_user(user_id).await?;
// Verify ownership before deleting let found = devices.iter().any(|d| d.id() == device_id);
let devices = self.device_code_storage.list_by_user(user_id).await?; if !found {
let found = devices.iter().any(|d| d.id() == device_id); return Err(DomainError::new(
if !found { ErrorKind::NotFound,
return Err(DomainError::new( "DeviceCode",
ErrorKind::NotFound, "Device authorization not found or not owned by you",
"DeviceCode", ));
"Device authorization not found or not owned by you", }
)); self.device_code_storage.delete_by_id(device_id).await
} }
self.device_code_storage.delete_by_id(device_id).await }
}
} // ============================================================================
// Poll error (typed for RFC 8628 error responses)
// ============================================================================ // ============================================================================
// Poll error (typed for RFC 8628 error responses)
// ============================================================================ /// Typed errors for the device token polling endpoint (RFC 8628 §3.5).
#[derive(Debug)]
/// Typed errors for the device token polling endpoint (RFC 8628 §3.5). pub enum DevicePollError {
#[derive(Debug)] /// The authorization request is still pending (user hasn't acted yet).
pub enum DevicePollError { AuthorizationPending,
/// The authorization request is still pending (user hasn't acted yet). /// The client is polling too fast; increase the interval.
AuthorizationPending, SlowDown,
/// The client is polling too fast; increase the interval. /// The user denied the authorization request.
SlowDown, AccessDenied,
/// The user denied the authorization request. /// The device_code has expired.
AccessDenied, ExpiredToken,
/// The device_code has expired. /// The device_code is not recognized.
ExpiredToken, InvalidDeviceCode,
/// The device_code is not recognized. }
InvalidDeviceCode,
} impl DevicePollError {
/// RFC 8628 error string for the JSON response.
impl DevicePollError { pub fn error_code(&self) -> &'static str {
/// RFC 8628 error string for the JSON response. match self {
pub fn error_code(&self) -> &'static str { Self::AuthorizationPending => "authorization_pending",
match self { Self::SlowDown => "slow_down",
Self::AuthorizationPending => "authorization_pending", Self::AccessDenied => "access_denied",
Self::SlowDown => "slow_down", Self::ExpiredToken => "expired_token",
Self::AccessDenied => "access_denied", Self::InvalidDeviceCode => "invalid_grant",
Self::ExpiredToken => "expired_token", }
Self::InvalidDeviceCode => "invalid_grant", }
}
} pub fn description(&self) -> &'static str {
match self {
pub fn description(&self) -> &'static str { Self::AuthorizationPending => {
match self { "The authorization request is still pending. Continue polling."
Self::AuthorizationPending => { }
"The authorization request is still pending. Continue polling." Self::SlowDown => "You are polling too frequently. Please slow down.",
} Self::AccessDenied => "The user denied the authorization request.",
Self::SlowDown => "You are polling too frequently. Please slow down.", Self::ExpiredToken => {
Self::AccessDenied => "The user denied the authorization request.", "The device_code has expired. Please start a new authorization flow."
Self::ExpiredToken => { }
"The device_code has expired. Please start a new authorization flow." Self::InvalidDeviceCode => "The device_code is not recognized.",
} }
Self::InvalidDeviceCode => "The device_code is not recognized.", }
}
} /// HTTP status code per RFC 8628 §3.5:
/// - authorization_pending and slow_down: 400
/// HTTP status code per RFC 8628 §3.5: /// - access_denied: 403
/// - authorization_pending and slow_down: 400 /// - expired_token: 400
/// - access_denied: 403 pub fn http_status(&self) -> u16 {
/// - expired_token: 400 match self {
pub fn http_status(&self) -> u16 { Self::AuthorizationPending | Self::SlowDown | Self::ExpiredToken => 400,
match self { Self::AccessDenied => 403,
Self::AuthorizationPending | Self::SlowDown | Self::ExpiredToken => 400, Self::InvalidDeviceCode => 400,
Self::AccessDenied => 403, }
Self::InvalidDeviceCode => 400, }
} }
}
} // ============================================================================
// Additional DTOs (used by service, not in the handler module)
// ============================================================================ // ============================================================================
// Additional DTOs (used by service, not in the handler module)
// ============================================================================ /// DTO for listing authorized devices in the user's profile.
#[derive(Debug, serde::Serialize)]
/// DTO for listing authorized devices in the user's profile. pub struct DeviceInfoDto {
#[derive(Debug, serde::Serialize)] pub id: String,
pub struct DeviceInfoDto { pub client_name: String,
pub id: String, pub scopes: String,
pub client_name: String, pub status: String,
pub scopes: String, pub created_at: String,
pub status: String, pub authorized_at: Option<String>,
pub created_at: String, pub expires_at: String,
pub authorized_at: Option<String>, }
pub expires_at: String,
} // ============================================================================
// Helpers
// ============================================================================ // ============================================================================
// Helpers
// ============================================================================ /// Generate a cryptographically random device_code (hex-encoded).
fn generate_device_code() -> String {
/// Generate a cryptographically random device_code (hex-encoded). use rand_core::{OsRng, RngCore};
fn generate_device_code() -> String { let mut bytes = [0u8; DEVICE_CODE_BYTES];
use rand_core::{OsRng, RngCore}; OsRng.fill_bytes(&mut bytes);
let mut bytes = [0u8; DEVICE_CODE_BYTES]; hex::encode(bytes)
OsRng.fill_bytes(&mut bytes); }
hex::encode(bytes)
} /// Generate a human-readable user_code in the format "ABCD-1234".
fn generate_user_code() -> String {
/// Generate a human-readable user_code in the format "ABCD-1234". use rand_core::{OsRng, RngCore};
fn generate_user_code() -> String { let mut rng_bytes = [0u8; 8];
use rand_core::{OsRng, RngCore}; OsRng.fill_bytes(&mut rng_bytes);
let mut rng_bytes = [0u8; 8];
OsRng.fill_bytes(&mut rng_bytes); let letters: String = (0..USER_CODE_LETTER_LEN)
.map(|i| {
let letters: String = (0..USER_CODE_LETTER_LEN) let b = rng_bytes[i] % 26;
.map(|i| { (b'A' + b) as char
let b = rng_bytes[i] % 26; })
(b'A' + b) as char .collect();
})
.collect(); let digits: String = (0..USER_CODE_DIGIT_LEN)
.map(|i| {
let digits: String = (0..USER_CODE_DIGIT_LEN) let b = rng_bytes[USER_CODE_LETTER_LEN + i] % 10;
.map(|i| { (b'0' + b) as char
let b = rng_bytes[USER_CODE_LETTER_LEN + i] % 10; })
(b'0' + b) as char .collect();
})
.collect(); format!("{}-{}", letters, digits)
}
format!("{}-{}", letters, digits)
}
+20 -12
View File
@@ -398,12 +398,16 @@ impl FolderUseCase for FolderService {
} }
// Rename folder — UPDATE RETURNING gives us the updated row directly // Rename folder — UPDATE RETURNING gives us the updated row directly
let folder = self.folder_storage.rename_folder(id, dto.name).await.map_err(|e| { let folder = self
DomainError::internal_error( .folder_storage
"FolderStorage", .rename_folder(id, dto.name)
format!("Failed to rename folder with ID: {}: {}", id, e), .await
) .map_err(|e| {
})?; DomainError::internal_error(
"FolderStorage",
format!("Failed to rename folder with ID: {}: {}", id, e),
)
})?;
Ok(FolderDto::from(folder)) Ok(FolderDto::from(folder))
} }
@@ -455,12 +459,16 @@ impl FolderUseCase for FolderService {
// Move folder — UPDATE RETURNING gives us the updated row directly // Move folder — UPDATE RETURNING gives us the updated row directly
let parent_ref = dto.parent_id.as_deref(); let parent_ref = dto.parent_id.as_deref();
let folder = self.folder_storage.move_folder(id, parent_ref).await.map_err(|e| { let folder = self
DomainError::internal_error( .folder_storage
"FolderStorage", .move_folder(id, parent_ref)
format!("Failed to move folder with ID: {}: {}", id, e), .await
) .map_err(|e| {
})?; DomainError::internal_error(
"FolderStorage",
format!("Failed to move folder with ID: {}: {}", id, e),
)
})?;
Ok(FolderDto::from(folder)) Ok(FolderDto::from(folder))
} }
+1 -1
View File
@@ -3,8 +3,8 @@ pub mod app_password_service;
pub mod auth_application_service; pub mod auth_application_service;
pub mod batch_operations; pub mod batch_operations;
pub mod calendar_service; pub mod calendar_service;
pub mod device_auth_service;
pub mod contact_service; pub mod contact_service;
pub mod device_auth_service;
pub mod favorites_service; pub mod favorites_service;
pub mod file_management_service; pub mod file_management_service;
pub mod file_retrieval_service; pub mod file_retrieval_service;
+2 -1
View File
@@ -636,7 +636,8 @@ impl AppConfig {
{ {
config.auth.rate_limit.register_max_requests = val; config.auth.rate_limit.register_max_requests = val;
} }
if let Ok(v) = env::var("OXICLOUD_RATE_LIMIT_REGISTER_WINDOW_SECS").map(|v| v.parse::<u64>()) if let Ok(v) =
env::var("OXICLOUD_RATE_LIMIT_REGISTER_WINDOW_SECS").map(|v| v.parse::<u64>())
&& let Ok(val) = v && let Ok(val) = v
{ {
config.auth.rate_limit.register_window_secs = val; config.auth.rate_limit.register_window_secs = val;
+10 -7
View File
@@ -603,10 +603,11 @@ impl AppServiceFactory {
Arc::new(crate::infrastructure::repositories::UserPgRepository::new( Arc::new(crate::infrastructure::repositories::UserPgRepository::new(
pool.clone(), pool.clone(),
)); ));
let session_repo: Arc<dyn crate::application::ports::auth_ports::SessionStoragePort> = let session_repo: Arc<
Arc::new(crate::infrastructure::repositories::SessionPgRepository::new( dyn crate::application::ports::auth_ports::SessionStoragePort,
pool.clone(), > = Arc::new(
)); crate::infrastructure::repositories::SessionPgRepository::new(pool.clone()),
);
let base_url = self.config.base_url(); let base_url = self.config.base_url();
let device_auth_svc = Arc::new(DeviceAuthService::new( let device_auth_svc = Arc::new(DeviceAuthService::new(
@@ -625,8 +626,9 @@ impl AppServiceFactory {
use crate::application::services::app_password_service::AppPasswordService; use crate::application::services::app_password_service::AppPasswordService;
use crate::infrastructure::repositories::AppPasswordPgRepository; use crate::infrastructure::repositories::AppPasswordPgRepository;
let app_pw_repo: Arc<dyn crate::application::ports::auth_ports::AppPasswordStoragePort> = let app_pw_repo: Arc<
Arc::new(AppPasswordPgRepository::new(pool.clone())); dyn crate::application::ports::auth_ports::AppPasswordStoragePort,
> = Arc::new(AppPasswordPgRepository::new(pool.clone()));
let hasher: Arc<dyn crate::application::ports::auth_ports::PasswordHasherPort> = let hasher: Arc<dyn crate::application::ports::auth_ports::PasswordHasherPort> =
Arc::new( Arc::new(
crate::infrastructure::services::password_hasher::Argon2PasswordHasher::new( crate::infrastructure::services::password_hasher::Argon2PasswordHasher::new(
@@ -818,7 +820,8 @@ pub struct ApplicationServices {
pub struct AuthServices { pub struct AuthServices {
pub token_service: Arc<dyn crate::application::ports::auth_ports::TokenServicePort>, pub token_service: Arc<dyn crate::application::ports::auth_ports::TokenServicePort>,
pub auth_application_service: Arc<AuthApplicationService>, pub auth_application_service: Arc<AuthApplicationService>,
pub login_lockout: Arc<crate::infrastructure::services::login_lockout_service::LoginLockoutService>, pub login_lockout:
Arc<crate::infrastructure::services::login_lockout_service::LoginLockoutService>,
} }
/// Global application state for dependency injection /// Global application state for dependency injection
+262 -267
View File
@@ -1,267 +1,262 @@
//! Device Authorization Code entity (RFC 8628). //! Device Authorization Code entity (RFC 8628).
//! //!
//! Represents a pending or completed OAuth 2.0 Device Authorization Grant flow. //! Represents a pending or completed OAuth 2.0 Device Authorization Grant flow.
use chrono::{DateTime, Duration, Utc}; use chrono::{DateTime, Duration, Utc};
use uuid::Uuid; use uuid::Uuid;
/// Status of a device authorization flow. /// Status of a device authorization flow.
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DeviceCodeStatus { pub enum DeviceCodeStatus {
/// Waiting for the user to authorize on the verification page. /// Waiting for the user to authorize on the verification page.
Pending, Pending,
/// User approved — tokens are ready for the polling client. /// User approved — tokens are ready for the polling client.
Authorized, Authorized,
/// User explicitly denied the request. /// User explicitly denied the request.
Denied, Denied,
/// The code expired before the user acted. /// The code expired before the user acted.
Expired, Expired,
} }
impl DeviceCodeStatus { impl DeviceCodeStatus {
pub fn as_str(&self) -> &'static str { pub fn as_str(&self) -> &'static str {
match self { match self {
Self::Pending => "pending", Self::Pending => "pending",
Self::Authorized => "authorized", Self::Authorized => "authorized",
Self::Denied => "denied", Self::Denied => "denied",
Self::Expired => "expired", Self::Expired => "expired",
} }
} }
pub fn from_str(s: &str) -> Option<Self> { pub fn from_str(s: &str) -> Option<Self> {
match s { match s {
"pending" => Some(Self::Pending), "pending" => Some(Self::Pending),
"authorized" => Some(Self::Authorized), "authorized" => Some(Self::Authorized),
"denied" => Some(Self::Denied), "denied" => Some(Self::Denied),
"expired" => Some(Self::Expired), "expired" => Some(Self::Expired),
_ => None, _ => None,
} }
} }
} }
impl std::fmt::Display for DeviceCodeStatus { impl std::fmt::Display for DeviceCodeStatus {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.as_str()) write!(f, "{}", self.as_str())
} }
} }
/// Domain entity for a Device Authorization flow. /// Domain entity for a Device Authorization flow.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct DeviceCode { pub struct DeviceCode {
id: String, id: String,
device_code: String, device_code: String,
user_code: String, user_code: String,
client_name: String, client_name: String,
scopes: String, scopes: String,
status: DeviceCodeStatus, status: DeviceCodeStatus,
user_id: Option<String>, user_id: Option<String>,
access_token: Option<String>, access_token: Option<String>,
refresh_token: Option<String>, refresh_token: Option<String>,
verification_uri: String, verification_uri: String,
verification_uri_complete: Option<String>, verification_uri_complete: Option<String>,
expires_at: DateTime<Utc>, expires_at: DateTime<Utc>,
poll_interval_secs: i32, poll_interval_secs: i32,
last_poll_at: Option<DateTime<Utc>>, last_poll_at: Option<DateTime<Utc>>,
created_at: DateTime<Utc>, created_at: DateTime<Utc>,
authorized_at: Option<DateTime<Utc>>, authorized_at: Option<DateTime<Utc>>,
} }
impl DeviceCode { impl DeviceCode {
/// Create a new pending device code flow. /// Create a new pending device code flow.
/// ///
/// * `device_code` — opaque token for client polling (64 hex chars) /// * `device_code` — opaque token for client polling (64 hex chars)
/// * `user_code` — short human-readable code (e.g. "ABCD-1234") /// * `user_code` — short human-readable code (e.g. "ABCD-1234")
/// * `client_name` — display name of the requesting client /// * `client_name` — display name of the requesting client
/// * `scopes` — requested scopes (e.g. "webdav,caldav,carddav") /// * `scopes` — requested scopes (e.g. "webdav,caldav,carddav")
/// * `verification_uri` — URL the user must visit /// * `verification_uri` — URL the user must visit
/// * `expires_in_secs` — TTL for the device code /// * `expires_in_secs` — TTL for the device code
/// * `poll_interval_secs` — minimum polling interval /// * `poll_interval_secs` — minimum polling interval
pub fn new( pub fn new(
device_code: String, device_code: String,
user_code: String, user_code: String,
client_name: String, client_name: String,
scopes: String, scopes: String,
verification_uri: String, verification_uri: String,
verification_uri_complete: Option<String>, verification_uri_complete: Option<String>,
expires_in_secs: i64, expires_in_secs: i64,
poll_interval_secs: i32, poll_interval_secs: i32,
) -> Self { ) -> Self {
let now = Utc::now(); let now = Utc::now();
Self { Self {
id: Uuid::new_v4().to_string(), id: Uuid::new_v4().to_string(),
device_code, device_code,
user_code, user_code,
client_name, client_name,
scopes, scopes,
status: DeviceCodeStatus::Pending, status: DeviceCodeStatus::Pending,
user_id: None, user_id: None,
access_token: None, access_token: None,
refresh_token: None, refresh_token: None,
verification_uri, verification_uri,
verification_uri_complete, verification_uri_complete,
expires_at: now + Duration::seconds(expires_in_secs), expires_at: now + Duration::seconds(expires_in_secs),
poll_interval_secs, poll_interval_secs,
last_poll_at: None, last_poll_at: None,
created_at: now, created_at: now,
authorized_at: None, authorized_at: None,
} }
} }
/// Reconstruct from database row. /// Reconstruct from database row.
#[allow(clippy::too_many_arguments)] #[allow(clippy::too_many_arguments)]
pub fn from_raw( pub fn from_raw(
id: String, id: String,
device_code: String, device_code: String,
user_code: String, user_code: String,
client_name: String, client_name: String,
scopes: String, scopes: String,
status: DeviceCodeStatus, status: DeviceCodeStatus,
user_id: Option<String>, user_id: Option<String>,
access_token: Option<String>, access_token: Option<String>,
refresh_token: Option<String>, refresh_token: Option<String>,
verification_uri: String, verification_uri: String,
verification_uri_complete: Option<String>, verification_uri_complete: Option<String>,
expires_at: DateTime<Utc>, expires_at: DateTime<Utc>,
poll_interval_secs: i32, poll_interval_secs: i32,
last_poll_at: Option<DateTime<Utc>>, last_poll_at: Option<DateTime<Utc>>,
created_at: DateTime<Utc>, created_at: DateTime<Utc>,
authorized_at: Option<DateTime<Utc>>, authorized_at: Option<DateTime<Utc>>,
) -> Self { ) -> Self {
Self { Self {
id, id,
device_code, device_code,
user_code, user_code,
client_name, client_name,
scopes, scopes,
status, status,
user_id, user_id,
access_token, access_token,
refresh_token, refresh_token,
verification_uri, verification_uri,
verification_uri_complete, verification_uri_complete,
expires_at, expires_at,
poll_interval_secs, poll_interval_secs,
last_poll_at, last_poll_at,
created_at, created_at,
authorized_at, authorized_at,
} }
} }
// ── Getters ────────────────────────────────────────────────── // ── Getters ──────────────────────────────────────────────────
pub fn id(&self) -> &str { pub fn id(&self) -> &str {
&self.id &self.id
} }
pub fn device_code(&self) -> &str { pub fn device_code(&self) -> &str {
&self.device_code &self.device_code
} }
pub fn user_code(&self) -> &str { pub fn user_code(&self) -> &str {
&self.user_code &self.user_code
} }
pub fn client_name(&self) -> &str { pub fn client_name(&self) -> &str {
&self.client_name &self.client_name
} }
pub fn scopes(&self) -> &str { pub fn scopes(&self) -> &str {
&self.scopes &self.scopes
} }
pub fn status(&self) -> DeviceCodeStatus { pub fn status(&self) -> DeviceCodeStatus {
self.status self.status
} }
pub fn user_id(&self) -> Option<&str> { pub fn user_id(&self) -> Option<&str> {
self.user_id.as_deref() self.user_id.as_deref()
} }
pub fn access_token(&self) -> Option<&str> { pub fn access_token(&self) -> Option<&str> {
self.access_token.as_deref() self.access_token.as_deref()
} }
pub fn refresh_token(&self) -> Option<&str> { pub fn refresh_token(&self) -> Option<&str> {
self.refresh_token.as_deref() self.refresh_token.as_deref()
} }
pub fn verification_uri(&self) -> &str { pub fn verification_uri(&self) -> &str {
&self.verification_uri &self.verification_uri
} }
pub fn verification_uri_complete(&self) -> Option<&str> { pub fn verification_uri_complete(&self) -> Option<&str> {
self.verification_uri_complete.as_deref() self.verification_uri_complete.as_deref()
} }
pub fn expires_at(&self) -> DateTime<Utc> { pub fn expires_at(&self) -> DateTime<Utc> {
self.expires_at self.expires_at
} }
pub fn poll_interval_secs(&self) -> i32 { pub fn poll_interval_secs(&self) -> i32 {
self.poll_interval_secs self.poll_interval_secs
} }
pub fn last_poll_at(&self) -> Option<DateTime<Utc>> { pub fn last_poll_at(&self) -> Option<DateTime<Utc>> {
self.last_poll_at self.last_poll_at
} }
pub fn created_at(&self) -> DateTime<Utc> { pub fn created_at(&self) -> DateTime<Utc> {
self.created_at self.created_at
} }
pub fn authorized_at(&self) -> Option<DateTime<Utc>> { pub fn authorized_at(&self) -> Option<DateTime<Utc>> {
self.authorized_at self.authorized_at
} }
// ── Business logic ─────────────────────────────────────────── // ── Business logic ───────────────────────────────────────────
/// Whether the device code has expired. /// Whether the device code has expired.
pub fn is_expired(&self) -> bool { pub fn is_expired(&self) -> bool {
Utc::now() > self.expires_at Utc::now() > self.expires_at
} }
/// Seconds remaining until expiry (clamped to 0). /// Seconds remaining until expiry (clamped to 0).
pub fn seconds_remaining(&self) -> i64 { pub fn seconds_remaining(&self) -> i64 {
let remaining = (self.expires_at - Utc::now()).num_seconds(); let remaining = (self.expires_at - Utc::now()).num_seconds();
remaining.max(0) remaining.max(0)
} }
/// Whether the client is polling too fast (within poll_interval_secs). /// Whether the client is polling too fast (within poll_interval_secs).
pub fn is_polling_too_fast(&self) -> bool { pub fn is_polling_too_fast(&self) -> bool {
if let Some(last) = self.last_poll_at { if let Some(last) = self.last_poll_at {
let elapsed = (Utc::now() - last).num_seconds(); let elapsed = (Utc::now() - last).num_seconds();
elapsed < self.poll_interval_secs as i64 elapsed < self.poll_interval_secs as i64
} else { } else {
false false
} }
} }
/// Record a poll attempt timestamp. /// Record a poll attempt timestamp.
pub fn record_poll(&mut self) { pub fn record_poll(&mut self) {
self.last_poll_at = Some(Utc::now()); self.last_poll_at = Some(Utc::now());
} }
/// Authorize this device code for a specific user, storing the tokens. /// Authorize this device code for a specific user, storing the tokens.
pub fn authorize( pub fn authorize(&mut self, user_id: String, access_token: String, refresh_token: String) {
&mut self, self.status = DeviceCodeStatus::Authorized;
user_id: String, self.user_id = Some(user_id);
access_token: String, self.access_token = Some(access_token);
refresh_token: String, self.refresh_token = Some(refresh_token);
) { self.authorized_at = Some(Utc::now());
self.status = DeviceCodeStatus::Authorized; }
self.user_id = Some(user_id);
self.access_token = Some(access_token); /// Deny this device code.
self.refresh_token = Some(refresh_token); pub fn deny(&mut self) {
self.authorized_at = Some(Utc::now()); self.status = DeviceCodeStatus::Denied;
} }
/// Deny this device code. /// Mark as expired.
pub fn deny(&mut self) { pub fn mark_expired(&mut self) {
self.status = DeviceCodeStatus::Denied; self.status = DeviceCodeStatus::Expired;
} }
}
/// Mark as expired.
pub fn mark_expired(&mut self) {
self.status = DeviceCodeStatus::Expired;
}
}
+3 -2
View File
@@ -3,6 +3,7 @@ pub mod pg;
// Re-exportar para facilitar acceso // Re-exportar para facilitar acceso
pub use pg::{ pub use pg::{
AppPasswordPgRepository, DeviceCodePgRepository, FileBlobReadRepository, FileBlobWriteRepository, AppPasswordPgRepository, DeviceCodePgRepository, FileBlobReadRepository,
FolderDbRepository, SessionPgRepository, TrashDbRepository, UserPgRepository, FileBlobWriteRepository, FolderDbRepository, SessionPgRepository, TrashDbRepository,
UserPgRepository,
}; };
@@ -1,188 +1,178 @@
//! PostgreSQL repository for App Passwords. //! PostgreSQL repository for App Passwords.
use crate::application::ports::auth_ports::AppPasswordStoragePort; use crate::application::ports::auth_ports::AppPasswordStoragePort;
use crate::common::errors::DomainError; use crate::common::errors::DomainError;
use crate::domain::entities::app_password::AppPassword; use crate::domain::entities::app_password::AppPassword;
use async_trait::async_trait; use async_trait::async_trait;
use chrono::{DateTime, Utc}; use chrono::{DateTime, Utc};
use sqlx::PgPool; use sqlx::PgPool;
use std::sync::Arc; use std::sync::Arc;
pub struct AppPasswordPgRepository { pub struct AppPasswordPgRepository {
pool: Arc<PgPool>, pool: Arc<PgPool>,
} }
impl AppPasswordPgRepository { impl AppPasswordPgRepository {
pub fn new(pool: Arc<PgPool>) -> Self { pub fn new(pool: Arc<PgPool>) -> Self {
Self { pool } Self { pool }
} }
fn pool(&self) -> &PgPool { fn pool(&self) -> &PgPool {
&self.pool &self.pool
} }
} }
#[async_trait] #[async_trait]
impl AppPasswordStoragePort for AppPasswordPgRepository { impl AppPasswordStoragePort for AppPasswordPgRepository {
async fn create(&self, ap: AppPassword) -> Result<AppPassword, DomainError> { async fn create(&self, ap: AppPassword) -> Result<AppPassword, DomainError> {
sqlx::query( sqlx::query(
r#" r#"
INSERT INTO auth.app_passwords INSERT INTO auth.app_passwords
(id, user_id, label, password_hash, prefix, scopes, (id, user_id, label, password_hash, prefix, scopes,
created_at, last_used_at, expires_at, active) created_at, last_used_at, expires_at, active)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
"#, "#,
) )
.bind(&ap.id) .bind(&ap.id)
.bind(&ap.user_id) .bind(&ap.user_id)
.bind(&ap.label) .bind(&ap.label)
.bind(&ap.password_hash) .bind(&ap.password_hash)
.bind(&ap.prefix) .bind(&ap.prefix)
.bind(&ap.scopes) .bind(&ap.scopes)
.bind(ap.created_at) .bind(ap.created_at)
.bind(ap.last_used_at) .bind(ap.last_used_at)
.bind(ap.expires_at) .bind(ap.expires_at)
.bind(ap.active) .bind(ap.active)
.execute(self.pool()) .execute(self.pool())
.await .await
.map_err(|e| DomainError::internal_error("AppPasswordPg", format!("create: {e}")))?; .map_err(|e| DomainError::internal_error("AppPasswordPg", format!("create: {e}")))?;
Ok(ap) Ok(ap)
} }
async fn list_by_user(&self, user_id: &str) -> Result<Vec<AppPassword>, DomainError> { async fn list_by_user(&self, user_id: &str) -> Result<Vec<AppPassword>, DomainError> {
let rows = sqlx::query_as::<_, AppPasswordRow>( let rows = sqlx::query_as::<_, AppPasswordRow>(
r#" r#"
SELECT id, user_id, label, password_hash, prefix, scopes, SELECT id, user_id, label, password_hash, prefix, scopes,
created_at, last_used_at, expires_at, active created_at, last_used_at, expires_at, active
FROM auth.app_passwords FROM auth.app_passwords
WHERE user_id = $1 WHERE user_id = $1
ORDER BY created_at DESC ORDER BY created_at DESC
"#, "#,
) )
.bind(user_id) .bind(user_id)
.fetch_all(self.pool()) .fetch_all(self.pool())
.await .await
.map_err(|e| DomainError::internal_error("AppPasswordPg", format!("list: {e}")))?; .map_err(|e| DomainError::internal_error("AppPasswordPg", format!("list: {e}")))?;
Ok(rows.into_iter().map(|r| r.into()).collect()) Ok(rows.into_iter().map(|r| r.into()).collect())
} }
async fn get_by_id(&self, id: &str) -> Result<AppPassword, DomainError> { async fn get_by_id(&self, id: &str) -> Result<AppPassword, DomainError> {
let row = sqlx::query_as::<_, AppPasswordRow>( let row = sqlx::query_as::<_, AppPasswordRow>(
r#" r#"
SELECT id, user_id, label, password_hash, prefix, scopes, SELECT id, user_id, label, password_hash, prefix, scopes,
created_at, last_used_at, expires_at, active created_at, last_used_at, expires_at, active
FROM auth.app_passwords FROM auth.app_passwords
WHERE id = $1 WHERE id = $1
"#, "#,
) )
.bind(id) .bind(id)
.fetch_optional(self.pool()) .fetch_optional(self.pool())
.await .await
.map_err(|e| DomainError::internal_error("AppPasswordPg", format!("get_by_id: {e}")))? .map_err(|e| DomainError::internal_error("AppPasswordPg", format!("get_by_id: {e}")))?
.ok_or_else(|| DomainError::not_found("AppPassword", id))?; .ok_or_else(|| DomainError::not_found("AppPassword", id))?;
Ok(row.into()) Ok(row.into())
} }
async fn get_active_by_user_id( async fn get_active_by_user_id(&self, user_id: &str) -> Result<Vec<AppPassword>, DomainError> {
&self, let rows = sqlx::query_as::<_, AppPasswordRow>(
user_id: &str, r#"
) -> Result<Vec<AppPassword>, DomainError> { SELECT id, user_id, label, password_hash, prefix, scopes,
let rows = sqlx::query_as::<_, AppPasswordRow>( created_at, last_used_at, expires_at, active
r#" FROM auth.app_passwords
SELECT id, user_id, label, password_hash, prefix, scopes, WHERE user_id = $1
created_at, last_used_at, expires_at, active AND active = TRUE
FROM auth.app_passwords AND (expires_at IS NULL OR expires_at > NOW())
WHERE user_id = $1 "#,
AND active = TRUE )
AND (expires_at IS NULL OR expires_at > NOW()) .bind(user_id)
"#, .fetch_all(self.pool())
) .await
.bind(user_id) .map_err(|e| DomainError::internal_error("AppPasswordPg", format!("get_active: {e}")))?;
.fetch_all(self.pool())
.await Ok(rows.into_iter().map(|r| r.into()).collect())
.map_err(|e| { }
DomainError::internal_error("AppPasswordPg", format!("get_active: {e}"))
})?; async fn touch_last_used(&self, id: &str) -> Result<(), DomainError> {
sqlx::query("UPDATE auth.app_passwords SET last_used_at = NOW() WHERE id = $1")
Ok(rows.into_iter().map(|r| r.into()).collect()) .bind(id)
} .execute(self.pool())
.await
async fn touch_last_used(&self, id: &str) -> Result<(), DomainError> { .map_err(|e| DomainError::internal_error("AppPasswordPg", format!("touch: {e}")))?;
sqlx::query("UPDATE auth.app_passwords SET last_used_at = NOW() WHERE id = $1") Ok(())
.bind(id) }
.execute(self.pool())
.await async fn revoke(&self, id: &str) -> Result<(), DomainError> {
.map_err(|e| { let result = sqlx::query("UPDATE auth.app_passwords SET active = FALSE WHERE id = $1")
DomainError::internal_error("AppPasswordPg", format!("touch: {e}")) .bind(id)
})?; .execute(self.pool())
Ok(()) .await
} .map_err(|e| DomainError::internal_error("AppPasswordPg", format!("revoke: {e}")))?;
async fn revoke(&self, id: &str) -> Result<(), DomainError> { if result.rows_affected() == 0 {
let result = return Err(DomainError::not_found("AppPassword", id));
sqlx::query("UPDATE auth.app_passwords SET active = FALSE WHERE id = $1") }
.bind(id) Ok(())
.execute(self.pool()) }
.await
.map_err(|e| { async fn delete_expired(&self) -> Result<u64, DomainError> {
DomainError::internal_error("AppPasswordPg", format!("revoke: {e}")) let result = sqlx::query(
})?; r#"
DELETE FROM auth.app_passwords
if result.rows_affected() == 0 { WHERE (active = FALSE)
return Err(DomainError::not_found("AppPassword", id)); OR (expires_at IS NOT NULL AND expires_at < NOW())
} "#,
Ok(()) )
} .execute(self.pool())
.await
async fn delete_expired(&self) -> Result<u64, DomainError> { .map_err(|e| {
let result = sqlx::query( DomainError::internal_error("AppPasswordPg", format!("delete_expired: {e}"))
r#" })?;
DELETE FROM auth.app_passwords
WHERE (active = FALSE) Ok(result.rows_affected())
OR (expires_at IS NOT NULL AND expires_at < NOW()) }
"#, }
)
.execute(self.pool()) /// Internal row struct for sqlx mapping.
.await #[derive(sqlx::FromRow)]
.map_err(|e| { struct AppPasswordRow {
DomainError::internal_error("AppPasswordPg", format!("delete_expired: {e}")) id: String,
})?; user_id: String,
label: String,
Ok(result.rows_affected()) password_hash: String,
} prefix: String,
} scopes: String,
created_at: DateTime<Utc>,
/// Internal row struct for sqlx mapping. last_used_at: Option<DateTime<Utc>>,
#[derive(sqlx::FromRow)] expires_at: Option<DateTime<Utc>>,
struct AppPasswordRow { active: bool,
id: String, }
user_id: String,
label: String, impl From<AppPasswordRow> for AppPassword {
password_hash: String, fn from(r: AppPasswordRow) -> Self {
prefix: String, AppPassword {
scopes: String, id: r.id,
created_at: DateTime<Utc>, user_id: r.user_id,
last_used_at: Option<DateTime<Utc>>, label: r.label,
expires_at: Option<DateTime<Utc>>, password_hash: r.password_hash,
active: bool, prefix: r.prefix,
} scopes: r.scopes,
created_at: r.created_at,
impl From<AppPasswordRow> for AppPassword { last_used_at: r.last_used_at,
fn from(r: AppPasswordRow) -> Self { expires_at: r.expires_at,
AppPassword { active: r.active,
id: r.id, }
user_id: r.user_id, }
label: r.label, }
password_hash: r.password_hash,
prefix: r.prefix,
scopes: r.scopes,
created_at: r.created_at,
last_used_at: r.last_used_at,
expires_at: r.expires_at,
active: r.active,
}
}
}
@@ -1,261 +1,259 @@
//! PostgreSQL repository for Device Authorization Grant (RFC 8628) codes. //! PostgreSQL repository for Device Authorization Grant (RFC 8628) codes.
use async_trait::async_trait; use async_trait::async_trait;
use sqlx::{PgPool, Row}; use sqlx::{PgPool, Row};
use std::sync::Arc; use std::sync::Arc;
use crate::application::ports::auth_ports::DeviceCodeStoragePort; use crate::application::ports::auth_ports::DeviceCodeStoragePort;
use crate::common::errors::{DomainError, ErrorKind}; use crate::common::errors::{DomainError, ErrorKind};
use crate::domain::entities::device_code::{DeviceCode, DeviceCodeStatus}; use crate::domain::entities::device_code::{DeviceCode, DeviceCodeStatus};
pub struct DeviceCodePgRepository { pub struct DeviceCodePgRepository {
pool: Arc<PgPool>, pool: Arc<PgPool>,
} }
impl DeviceCodePgRepository { impl DeviceCodePgRepository {
pub fn new(pool: Arc<PgPool>) -> Self { pub fn new(pool: Arc<PgPool>) -> Self {
Self { pool } Self { pool }
} }
fn map_row(row: &sqlx::postgres::PgRow) -> Result<DeviceCode, DomainError> { fn map_row(row: &sqlx::postgres::PgRow) -> Result<DeviceCode, DomainError> {
let status_str: String = row.try_get("status").map_err(|e| { let status_str: String = row.try_get("status").map_err(|e| {
DomainError::new( DomainError::new(
ErrorKind::DatabaseError, ErrorKind::DatabaseError,
"DeviceCode", "DeviceCode",
format!("Failed to read status: {}", e), format!("Failed to read status: {}", e),
) )
})?; })?;
let status = DeviceCodeStatus::from_str(&status_str).unwrap_or(DeviceCodeStatus::Expired); let status = DeviceCodeStatus::from_str(&status_str).unwrap_or(DeviceCodeStatus::Expired);
Ok(DeviceCode::from_raw( Ok(DeviceCode::from_raw(
row.try_get("id").unwrap_or_default(), row.try_get("id").unwrap_or_default(),
row.try_get("device_code").unwrap_or_default(), row.try_get("device_code").unwrap_or_default(),
row.try_get("user_code").unwrap_or_default(), row.try_get("user_code").unwrap_or_default(),
row.try_get("client_name").unwrap_or_default(), row.try_get("client_name").unwrap_or_default(),
row.try_get("scopes").unwrap_or_default(), row.try_get("scopes").unwrap_or_default(),
status, status,
row.try_get("user_id").ok(), row.try_get("user_id").ok(),
row.try_get("access_token").ok(), row.try_get("access_token").ok(),
row.try_get("refresh_token").ok(), row.try_get("refresh_token").ok(),
row.try_get("verification_uri").unwrap_or_default(), row.try_get("verification_uri").unwrap_or_default(),
row.try_get("verification_uri_complete").ok(), row.try_get("verification_uri_complete").ok(),
row.try_get("expires_at").unwrap_or_default(), row.try_get("expires_at").unwrap_or_default(),
row.try_get::<i32, _>("poll_interval_secs").unwrap_or(5), row.try_get::<i32, _>("poll_interval_secs").unwrap_or(5),
row.try_get("last_poll_at").ok(), row.try_get("last_poll_at").ok(),
row.try_get("created_at").unwrap_or_default(), row.try_get("created_at").unwrap_or_default(),
row.try_get("authorized_at").ok(), row.try_get("authorized_at").ok(),
)) ))
} }
} }
#[async_trait] #[async_trait]
impl DeviceCodeStoragePort for DeviceCodePgRepository { impl DeviceCodeStoragePort for DeviceCodePgRepository {
async fn create_device_code(&self, dc: DeviceCode) -> Result<DeviceCode, DomainError> { async fn create_device_code(&self, dc: DeviceCode) -> Result<DeviceCode, DomainError> {
sqlx::query( sqlx::query(
r#" r#"
INSERT INTO auth.device_codes ( INSERT INTO auth.device_codes (
id, device_code, user_code, client_name, scopes, status, id, device_code, user_code, client_name, scopes, status,
user_id, access_token, refresh_token, user_id, access_token, refresh_token,
verification_uri, verification_uri_complete, verification_uri, verification_uri_complete,
expires_at, poll_interval_secs, last_poll_at, expires_at, poll_interval_secs, last_poll_at,
created_at, authorized_at created_at, authorized_at
) VALUES ( ) VALUES (
$1, $2, $3, $4, $5, $6::auth.device_code_status, $1, $2, $3, $4, $5, $6::auth.device_code_status,
$7, $8, $9, $7, $8, $9,
$10, $11, $10, $11,
$12, $13, $14, $12, $13, $14,
$15, $16 $15, $16
) )
"#, "#,
) )
.bind(dc.id()) .bind(dc.id())
.bind(dc.device_code()) .bind(dc.device_code())
.bind(dc.user_code()) .bind(dc.user_code())
.bind(dc.client_name()) .bind(dc.client_name())
.bind(dc.scopes()) .bind(dc.scopes())
.bind(dc.status().as_str()) .bind(dc.status().as_str())
.bind(dc.user_id()) .bind(dc.user_id())
.bind(dc.access_token()) .bind(dc.access_token())
.bind(dc.refresh_token()) .bind(dc.refresh_token())
.bind(dc.verification_uri()) .bind(dc.verification_uri())
.bind(dc.verification_uri_complete()) .bind(dc.verification_uri_complete())
.bind(dc.expires_at()) .bind(dc.expires_at())
.bind(dc.poll_interval_secs()) .bind(dc.poll_interval_secs())
.bind(dc.last_poll_at()) .bind(dc.last_poll_at())
.bind(dc.created_at()) .bind(dc.created_at())
.bind(dc.authorized_at()) .bind(dc.authorized_at())
.execute(self.pool.as_ref()) .execute(self.pool.as_ref())
.await .await
.map_err(|e| { .map_err(|e| {
DomainError::new( DomainError::new(
ErrorKind::DatabaseError, ErrorKind::DatabaseError,
"DeviceCode", "DeviceCode",
format!("Failed to create device code: {}", e), format!("Failed to create device code: {}", e),
) )
})?; })?;
Ok(dc) Ok(dc)
} }
async fn get_by_device_code(&self, device_code: &str) -> Result<DeviceCode, DomainError> { async fn get_by_device_code(&self, device_code: &str) -> Result<DeviceCode, DomainError> {
let row = sqlx::query( let row = sqlx::query(
r#" r#"
SELECT id, device_code, user_code, client_name, scopes, SELECT id, device_code, user_code, client_name, scopes,
status::text AS status, user_id, access_token, refresh_token, status::text AS status, user_id, access_token, refresh_token,
verification_uri, verification_uri_complete, verification_uri, verification_uri_complete,
expires_at, poll_interval_secs, last_poll_at, expires_at, poll_interval_secs, last_poll_at,
created_at, authorized_at created_at, authorized_at
FROM auth.device_codes FROM auth.device_codes
WHERE device_code = $1 WHERE device_code = $1
"#, "#,
) )
.bind(device_code) .bind(device_code)
.fetch_one(self.pool.as_ref()) .fetch_one(self.pool.as_ref())
.await .await
.map_err(|e| match e { .map_err(|e| match e {
sqlx::Error::RowNotFound => DomainError::new( sqlx::Error::RowNotFound => {
ErrorKind::NotFound, DomainError::new(ErrorKind::NotFound, "DeviceCode", "Device code not found")
"DeviceCode", }
"Device code not found", _ => DomainError::new(
), ErrorKind::DatabaseError,
_ => DomainError::new( "DeviceCode",
ErrorKind::DatabaseError, format!("Failed to fetch device code: {}", e),
"DeviceCode", ),
format!("Failed to fetch device code: {}", e), })?;
),
})?; Self::map_row(&row)
}
Self::map_row(&row)
} async fn get_pending_by_user_code(&self, user_code: &str) -> Result<DeviceCode, DomainError> {
let row = sqlx::query(
async fn get_pending_by_user_code(&self, user_code: &str) -> Result<DeviceCode, DomainError> { r#"
let row = sqlx::query( SELECT id, device_code, user_code, client_name, scopes,
r#" status::text AS status, user_id, access_token, refresh_token,
SELECT id, device_code, user_code, client_name, scopes, verification_uri, verification_uri_complete,
status::text AS status, user_id, access_token, refresh_token, expires_at, poll_interval_secs, last_poll_at,
verification_uri, verification_uri_complete, created_at, authorized_at
expires_at, poll_interval_secs, last_poll_at, FROM auth.device_codes
created_at, authorized_at WHERE user_code = $1
FROM auth.device_codes AND status = 'pending'
WHERE user_code = $1 AND expires_at > NOW()
AND status = 'pending' "#,
AND expires_at > NOW() )
"#, .bind(user_code)
) .fetch_one(self.pool.as_ref())
.bind(user_code) .await
.fetch_one(self.pool.as_ref()) .map_err(|e| match e {
.await sqlx::Error::RowNotFound => DomainError::new(
.map_err(|e| match e { ErrorKind::NotFound,
sqlx::Error::RowNotFound => DomainError::new( "DeviceCode",
ErrorKind::NotFound, "User code not found or expired",
"DeviceCode", ),
"User code not found or expired", _ => DomainError::new(
), ErrorKind::DatabaseError,
_ => DomainError::new( "DeviceCode",
ErrorKind::DatabaseError, format!("Failed to fetch by user code: {}", e),
"DeviceCode", ),
format!("Failed to fetch by user code: {}", e), })?;
),
})?; Self::map_row(&row)
}
Self::map_row(&row)
} async fn update_device_code(&self, dc: DeviceCode) -> Result<(), DomainError> {
sqlx::query(
async fn update_device_code(&self, dc: DeviceCode) -> Result<(), DomainError> { r#"
sqlx::query( UPDATE auth.device_codes SET
r#" status = $2::auth.device_code_status,
UPDATE auth.device_codes SET user_id = $3,
status = $2::auth.device_code_status, access_token = $4,
user_id = $3, refresh_token = $5,
access_token = $4, last_poll_at = $6,
refresh_token = $5, authorized_at = $7
last_poll_at = $6, WHERE id = $1
authorized_at = $7 "#,
WHERE id = $1 )
"#, .bind(dc.id())
) .bind(dc.status().as_str())
.bind(dc.id()) .bind(dc.user_id())
.bind(dc.status().as_str()) .bind(dc.access_token())
.bind(dc.user_id()) .bind(dc.refresh_token())
.bind(dc.access_token()) .bind(dc.last_poll_at())
.bind(dc.refresh_token()) .bind(dc.authorized_at())
.bind(dc.last_poll_at()) .execute(self.pool.as_ref())
.bind(dc.authorized_at()) .await
.execute(self.pool.as_ref()) .map_err(|e| {
.await DomainError::new(
.map_err(|e| { ErrorKind::DatabaseError,
DomainError::new( "DeviceCode",
ErrorKind::DatabaseError, format!("Failed to update device code: {}", e),
"DeviceCode", )
format!("Failed to update device code: {}", e), })?;
)
})?; Ok(())
}
Ok(())
} async fn delete_expired(&self) -> Result<u64, DomainError> {
let result = sqlx::query(
async fn delete_expired(&self) -> Result<u64, DomainError> { r#"
let result = sqlx::query( DELETE FROM auth.device_codes
r#" WHERE expires_at < NOW()
DELETE FROM auth.device_codes AND status IN ('pending', 'expired')
WHERE expires_at < NOW() "#,
AND status IN ('pending', 'expired') )
"#, .execute(self.pool.as_ref())
) .await
.execute(self.pool.as_ref()) .map_err(|e| {
.await DomainError::new(
.map_err(|e| { ErrorKind::DatabaseError,
DomainError::new( "DeviceCode",
ErrorKind::DatabaseError, format!("Failed to delete expired device codes: {}", e),
"DeviceCode", )
format!("Failed to delete expired device codes: {}", e), })?;
)
})?; Ok(result.rows_affected())
}
Ok(result.rows_affected())
} async fn list_by_user(&self, user_id: &str) -> Result<Vec<DeviceCode>, DomainError> {
let rows = sqlx::query(
async fn list_by_user(&self, user_id: &str) -> Result<Vec<DeviceCode>, DomainError> { r#"
let rows = sqlx::query( SELECT id, device_code, user_code, client_name, scopes,
r#" status::text AS status, user_id, access_token, refresh_token,
SELECT id, device_code, user_code, client_name, scopes, verification_uri, verification_uri_complete,
status::text AS status, user_id, access_token, refresh_token, expires_at, poll_interval_secs, last_poll_at,
verification_uri, verification_uri_complete, created_at, authorized_at
expires_at, poll_interval_secs, last_poll_at, FROM auth.device_codes
created_at, authorized_at WHERE user_id = $1
FROM auth.device_codes ORDER BY created_at DESC
WHERE user_id = $1 "#,
ORDER BY created_at DESC )
"#, .bind(user_id)
) .fetch_all(self.pool.as_ref())
.bind(user_id) .await
.fetch_all(self.pool.as_ref()) .map_err(|e| {
.await DomainError::new(
.map_err(|e| { ErrorKind::DatabaseError,
DomainError::new( "DeviceCode",
ErrorKind::DatabaseError, format!("Failed to list device codes: {}", e),
"DeviceCode", )
format!("Failed to list device codes: {}", e), })?;
)
})?; rows.iter().map(Self::map_row).collect()
}
rows.iter().map(Self::map_row).collect()
} async fn delete_by_id(&self, id: &str) -> Result<(), DomainError> {
sqlx::query("DELETE FROM auth.device_codes WHERE id = $1")
async fn delete_by_id(&self, id: &str) -> Result<(), DomainError> { .bind(id)
sqlx::query("DELETE FROM auth.device_codes WHERE id = $1") .execute(self.pool.as_ref())
.bind(id) .await
.execute(self.pool.as_ref()) .map_err(|e| {
.await DomainError::new(
.map_err(|e| { ErrorKind::DatabaseError,
DomainError::new( "DeviceCode",
ErrorKind::DatabaseError, format!("Failed to delete device code: {}", e),
"DeviceCode", )
format!("Failed to delete device code: {}", e), })?;
)
})?; Ok(())
}
Ok(()) }
}
}
@@ -268,10 +268,18 @@ impl FolderRepository for FolderDbRepository {
limit: usize, limit: usize,
include_total: bool, include_total: bool,
) -> Result<(Vec<Folder>, Option<usize>), DomainError> { ) -> Result<(Vec<Folder>, Option<usize>), DomainError> {
let rows: Vec<(String, String, String, Option<String>, String, i64, i64, i64)> = let rows: Vec<(
if let Some(pid) = parent_id { String,
sqlx::query_as( String,
r#" String,
Option<String>,
String,
i64,
i64,
i64,
)> = if let Some(pid) = parent_id {
sqlx::query_as(
r#"
SELECT id::text, name, path, parent_id::text, user_id, SELECT id::text, name, path, parent_id::text, user_id,
EXTRACT(EPOCH FROM created_at)::bigint, EXTRACT(EPOCH FROM created_at)::bigint,
EXTRACT(EPOCH FROM updated_at)::bigint, EXTRACT(EPOCH FROM updated_at)::bigint,
@@ -281,15 +289,15 @@ impl FolderRepository for FolderDbRepository {
ORDER BY name ORDER BY name
LIMIT $2 OFFSET $3 LIMIT $2 OFFSET $3
"#, "#,
) )
.bind(pid) .bind(pid)
.bind(limit as i64) .bind(limit as i64)
.bind(offset as i64) .bind(offset as i64)
.fetch_all(self.pool()) .fetch_all(self.pool())
.await .await
} else { } else {
sqlx::query_as( sqlx::query_as(
r#" r#"
SELECT id::text, name, path, parent_id::text, user_id, SELECT id::text, name, path, parent_id::text, user_id,
EXTRACT(EPOCH FROM created_at)::bigint, EXTRACT(EPOCH FROM created_at)::bigint,
EXTRACT(EPOCH FROM updated_at)::bigint, EXTRACT(EPOCH FROM updated_at)::bigint,
@@ -299,13 +307,13 @@ impl FolderRepository for FolderDbRepository {
ORDER BY name ORDER BY name
LIMIT $1 OFFSET $2 LIMIT $1 OFFSET $2
"#, "#,
) )
.bind(limit as i64) .bind(limit as i64)
.bind(offset as i64) .bind(offset as i64)
.fetch_all(self.pool()) .fetch_all(self.pool())
.await .await
} }
.map_err(|e| DomainError::internal_error("FolderDb", format!("paginate: {e}")))?; .map_err(|e| DomainError::internal_error("FolderDb", format!("paginate: {e}")))?;
// total_count is identical in every row; 0 when the result set is empty. // total_count is identical in every row; 0 when the result set is empty.
let total = if include_total { let total = if include_total {
@@ -333,10 +341,18 @@ impl FolderRepository for FolderDbRepository {
limit: usize, limit: usize,
include_total: bool, include_total: bool,
) -> Result<(Vec<Folder>, Option<usize>), DomainError> { ) -> Result<(Vec<Folder>, Option<usize>), DomainError> {
let rows: Vec<(String, String, String, Option<String>, String, i64, i64, i64)> = let rows: Vec<(
if let Some(pid) = parent_id { String,
sqlx::query_as( String,
r#" String,
Option<String>,
String,
i64,
i64,
i64,
)> = if let Some(pid) = parent_id {
sqlx::query_as(
r#"
SELECT id::text, name, path, parent_id::text, user_id, SELECT id::text, name, path, parent_id::text, user_id,
EXTRACT(EPOCH FROM created_at)::bigint, EXTRACT(EPOCH FROM created_at)::bigint,
EXTRACT(EPOCH FROM updated_at)::bigint, EXTRACT(EPOCH FROM updated_at)::bigint,
@@ -346,16 +362,16 @@ impl FolderRepository for FolderDbRepository {
ORDER BY name ORDER BY name
LIMIT $3 OFFSET $4 LIMIT $3 OFFSET $4
"#, "#,
) )
.bind(pid) .bind(pid)
.bind(owner_id) .bind(owner_id)
.bind(limit as i64) .bind(limit as i64)
.bind(offset as i64) .bind(offset as i64)
.fetch_all(self.pool()) .fetch_all(self.pool())
.await .await
} else { } else {
sqlx::query_as( sqlx::query_as(
r#" r#"
SELECT id::text, name, path, parent_id::text, user_id, SELECT id::text, name, path, parent_id::text, user_id,
EXTRACT(EPOCH FROM created_at)::bigint, EXTRACT(EPOCH FROM created_at)::bigint,
EXTRACT(EPOCH FROM updated_at)::bigint, EXTRACT(EPOCH FROM updated_at)::bigint,
@@ -365,16 +381,14 @@ impl FolderRepository for FolderDbRepository {
ORDER BY name ORDER BY name
LIMIT $2 OFFSET $3 LIMIT $2 OFFSET $3
"#, "#,
) )
.bind(owner_id) .bind(owner_id)
.bind(limit as i64) .bind(limit as i64)
.bind(offset as i64) .bind(offset as i64)
.fetch_all(self.pool()) .fetch_all(self.pool())
.await .await
} }
.map_err(|e| { .map_err(|e| DomainError::internal_error("FolderDb", format!("paginate_by_owner: {e}")))?;
DomainError::internal_error("FolderDb", format!("paginate_by_owner: {e}"))
})?;
let total = if include_total { let total = if include_total {
Some(rows.first().map_or(0, |r| r.7) as usize) Some(rows.first().map_or(0, |r| r.7) as usize)
@@ -867,10 +881,7 @@ impl FolderRepository for FolderDbRepository {
user_id: &str, user_id: &str,
) -> Result<Vec<Folder>, DomainError> { ) -> Result<Vec<Folder>, DomainError> {
let (where_extra, name_pattern) = match name_contains { let (where_extra, name_pattern) = match name_contains {
Some(name) if name.len() >= 3 => ( Some(name) if name.len() >= 3 => (" AND fo.name ILIKE $3", Some(format!("%{}%", name))),
" AND fo.name ILIKE $3",
Some(format!("%{}%", name)),
),
_ => ("", None), _ => ("", None),
}; };
+3 -1
View File
@@ -335,7 +335,9 @@ mod tests {
let claims1 = service.validate_token(&token).expect("Should validate"); let claims1 = service.validate_token(&token).expect("Should validate");
// Second call: cache hit — skips HMAC, returns cloned claims // Second call: cache hit — skips HMAC, returns cloned claims
let claims2 = service.validate_token(&token).expect("Should validate from cache"); let claims2 = service
.validate_token(&token)
.expect("Should validate from cache");
assert_eq!(claims1.sub, claims2.sub); assert_eq!(claims1.sub, claims2.sub);
assert_eq!(claims1.username, claims2.username); assert_eq!(claims1.username, claims2.username);
@@ -1,153 +1,150 @@
//! Account lockout service — blocks login for an account after N consecutive //! Account lockout service — blocks login for an account after N consecutive
//! failed attempts. //! failed attempts.
//! //!
//! Uses a `moka` TTL cache so that: //! Uses a `moka` TTL cache so that:
//! * Failed-attempt counters automatically expire after the lockout window. //! * Failed-attempt counters automatically expire after the lockout window.
//! * No database writes are needed — this is **in-memory** and therefore //! * No database writes are needed — this is **in-memory** and therefore
//! per-instance. If OxiCloud is deployed behind a load balancer with //! per-instance. If OxiCloud is deployed behind a load balancer with
//! multiple replicas, a sticky-session or shared Redis store would be //! multiple replicas, a sticky-session or shared Redis store would be
//! needed for cross-instance coordination (out of scope for v1). //! needed for cross-instance coordination (out of scope for v1).
//! //!
//! Typical flow: //! Typical flow:
//! 1. **Before password verification** → call [`LoginLockoutService::check`]. //! 1. **Before password verification** → call [`LoginLockoutService::check`].
//! If the account is locked, return `403` immediately without touching //! If the account is locked, return `403` immediately without touching
//! Argon2 (saves CPU). //! Argon2 (saves CPU).
//! 2. **After failed verification** → call [`LoginLockoutService::record_failure`]. //! 2. **After failed verification** → call [`LoginLockoutService::record_failure`].
//! 3. **After successful login** → call [`LoginLockoutService::record_success`] //! 3. **After successful login** → call [`LoginLockoutService::record_success`]
//! to reset the counter. //! to reset the counter.
use moka::sync::Cache; use moka::sync::Cache;
use std::time::Duration; use std::time::Duration;
/// Tracks consecutive failures for a single username. /// Tracks consecutive failures for a single username.
#[derive(Clone, Debug)] #[derive(Clone, Debug)]
struct FailureRecord { struct FailureRecord {
/// Number of consecutive failed attempts. /// Number of consecutive failed attempts.
count: u32, count: u32,
} }
/// In-memory account lockout tracker. /// In-memory account lockout tracker.
#[derive(Clone)] #[derive(Clone)]
pub struct LoginLockoutService { pub struct LoginLockoutService {
/// Maps `username -> FailureRecord`. TTL = lockout window. /// Maps `username -> FailureRecord`. TTL = lockout window.
cache: Cache<String, FailureRecord>, cache: Cache<String, FailureRecord>,
/// Maximum consecutive failures before the account is temporarily locked. /// Maximum consecutive failures before the account is temporarily locked.
max_failures: u32, max_failures: u32,
/// How long the lockout lasts (seconds). /// How long the lockout lasts (seconds).
lockout_secs: u64, lockout_secs: u64,
} }
impl LoginLockoutService { impl LoginLockoutService {
/// Create a new lockout service. /// Create a new lockout service.
/// ///
/// * `max_failures` — e.g. `5` (lock after 5 bad passwords) /// * `max_failures` — e.g. `5` (lock after 5 bad passwords)
/// * `lockout_secs` — e.g. `900` (15-minute lockout) /// * `lockout_secs` — e.g. `900` (15-minute lockout)
/// * `max_accounts` — upper bound on tracked accounts (evicts LRU) /// * `max_accounts` — upper bound on tracked accounts (evicts LRU)
pub fn new(max_failures: u32, lockout_secs: u64, max_accounts: u64) -> Self { pub fn new(max_failures: u32, lockout_secs: u64, max_accounts: u64) -> Self {
let cache = Cache::builder() let cache = Cache::builder()
.time_to_live(Duration::from_secs(lockout_secs)) .time_to_live(Duration::from_secs(lockout_secs))
.max_capacity(max_accounts) .max_capacity(max_accounts)
.build(); .build();
Self { Self {
cache, cache,
max_failures, max_failures,
lockout_secs, lockout_secs,
} }
} }
/// Check whether the account is currently locked. /// Check whether the account is currently locked.
/// ///
/// Returns `Ok(())` if the user may attempt login, or /// Returns `Ok(())` if the user may attempt login, or
/// `Err(remaining_secs)` with the *approximate* remaining lockout time. /// `Err(remaining_secs)` with the *approximate* remaining lockout time.
pub fn check(&self, username: &str) -> Result<(), u64> { pub fn check(&self, username: &str) -> Result<(), u64> {
if let Some(rec) = self.cache.get(&username.to_lowercase()) { if let Some(rec) = self.cache.get(&username.to_lowercase()) {
if rec.count >= self.max_failures { if rec.count >= self.max_failures {
// The entry exists and is over the threshold. Because moka // The entry exists and is over the threshold. Because moka
// evicts at TTL we know the lockout window has not yet elapsed. // evicts at TTL we know the lockout window has not yet elapsed.
return Err(self.lockout_secs); return Err(self.lockout_secs);
} }
} }
Ok(()) Ok(())
} }
/// Record a failed login attempt. Returns the new failure count. /// Record a failed login attempt. Returns the new failure count.
pub fn record_failure(&self, username: &str) -> u32 { pub fn record_failure(&self, username: &str) -> u32 {
let key = username.to_lowercase(); let key = username.to_lowercase();
let new_count = self let new_count = self.cache.get(&key).map(|r| r.count + 1).unwrap_or(1);
.cache self.cache
.get(&key) .insert(key.clone(), FailureRecord { count: new_count });
.map(|r| r.count + 1)
.unwrap_or(1); if new_count >= self.max_failures {
self.cache.insert(key.clone(), FailureRecord { count: new_count }); tracing::warn!(
username = %username,
if new_count >= self.max_failures { attempts = new_count,
tracing::warn!( lockout_secs = self.lockout_secs,
username = %username, "Account temporarily locked after {} consecutive failed login attempts",
attempts = new_count, new_count,
lockout_secs = self.lockout_secs, );
"Account temporarily locked after {} consecutive failed login attempts", }
new_count, new_count
); }
}
new_count /// Record a successful login — resets the failure counter.
} pub fn record_success(&self, username: &str) {
self.cache.invalidate(&username.to_lowercase());
/// Record a successful login — resets the failure counter. }
pub fn record_success(&self, username: &str) {
self.cache.invalidate(&username.to_lowercase()); /// Maximum failures before lockout (used to inform callers / error messages).
} pub fn max_failures(&self) -> u32 {
self.max_failures
/// Maximum failures before lockout (used to inform callers / error messages). }
pub fn max_failures(&self) -> u32 {
self.max_failures /// Lockout duration in seconds.
} pub fn lockout_secs(&self) -> u64 {
self.lockout_secs
/// Lockout duration in seconds. }
pub fn lockout_secs(&self) -> u64 { }
self.lockout_secs
} #[cfg(test)]
} mod tests {
use super::*;
#[cfg(test)]
mod tests { #[test]
use super::*; fn allows_login_under_threshold() {
let svc = LoginLockoutService::new(3, 60, 100);
#[test] assert!(svc.check("alice").is_ok());
fn allows_login_under_threshold() { svc.record_failure("alice");
let svc = LoginLockoutService::new(3, 60, 100); svc.record_failure("alice");
assert!(svc.check("alice").is_ok()); // 2 failures — still under threshold
svc.record_failure("alice"); assert!(svc.check("alice").is_ok());
svc.record_failure("alice"); }
// 2 failures — still under threshold
assert!(svc.check("alice").is_ok()); #[test]
} fn locks_after_threshold() {
let svc = LoginLockoutService::new(3, 60, 100);
#[test] svc.record_failure("bob");
fn locks_after_threshold() { svc.record_failure("bob");
let svc = LoginLockoutService::new(3, 60, 100); svc.record_failure("bob");
svc.record_failure("bob"); assert!(svc.check("bob").is_err());
svc.record_failure("bob"); }
svc.record_failure("bob");
assert!(svc.check("bob").is_err()); #[test]
} fn resets_on_success() {
let svc = LoginLockoutService::new(3, 60, 100);
#[test] svc.record_failure("carol");
fn resets_on_success() { svc.record_failure("carol");
let svc = LoginLockoutService::new(3, 60, 100); svc.record_success("carol");
svc.record_failure("carol"); // Counter reset — should be allowed again
svc.record_failure("carol"); assert!(svc.check("carol").is_ok());
svc.record_success("carol"); svc.record_failure("carol"); // starts over at 1
// Counter reset — should be allowed again assert!(svc.check("carol").is_ok());
assert!(svc.check("carol").is_ok()); }
svc.record_failure("carol"); // starts over at 1
assert!(svc.check("carol").is_ok()); #[test]
} fn case_insensitive() {
let svc = LoginLockoutService::new(2, 60, 100);
#[test] svc.record_failure("Dave");
fn case_insensitive() { svc.record_failure("dave");
let svc = LoginLockoutService::new(2, 60, 100); assert!(svc.check("DAVE").is_err());
svc.record_failure("Dave"); }
svc.record_failure("dave"); }
assert!(svc.check("DAVE").is_err());
}
}
+1 -1
View File
@@ -1,11 +1,11 @@
pub mod chunked_upload_service; pub mod chunked_upload_service;
pub mod compression_service; pub mod compression_service;
pub mod dedup_service; pub mod dedup_service;
pub mod login_lockout_service;
pub mod file_content_cache; pub mod file_content_cache;
pub mod file_system_i18n_service; pub mod file_system_i18n_service;
pub mod image_transcode_service; pub mod image_transcode_service;
pub mod jwt_service; pub mod jwt_service;
pub mod login_lockout_service;
pub mod oidc_service; pub mod oidc_service;
pub mod password_hasher; pub mod password_hasher;
pub mod path_resolver_service; pub mod path_resolver_service;
@@ -1,205 +1,219 @@
//! Single-query WebDAV path resolver. //! Single-query WebDAV path resolver.
//! //!
//! Replaces the double-query pattern (`get_folder_by_path` + `get_file_by_path`) //! Replaces the double-query pattern (`get_folder_by_path` + `get_file_by_path`)
//! with a single `UNION ALL` query that returns the first match. PostgreSQL's //! with a single `UNION ALL` query that returns the first match. PostgreSQL's
//! `Append` node short-circuits on `LIMIT 1`, so if the folder branch matches //! `Append` node short-circuits on `LIMIT 1`, so if the folder branch matches
//! the file branch is never executed. //! the file branch is never executed.
use sqlx::PgPool; use sqlx::PgPool;
use std::sync::Arc; use std::sync::Arc;
use crate::application::dtos::display_helpers::{ use crate::application::dtos::display_helpers::{
category_for, format_file_size, icon_class_for, icon_special_class_for, category_for, format_file_size, icon_class_for, icon_special_class_for,
}; };
use crate::application::dtos::file_dto::FileDto; use crate::application::dtos::file_dto::FileDto;
use crate::application::dtos::folder_dto::FolderDto; use crate::application::dtos::folder_dto::FolderDto;
use crate::common::errors::DomainError; use crate::common::errors::DomainError;
/// Result of resolving a WebDAV path — either a folder or a file. /// Result of resolving a WebDAV path — either a folder or a file.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub enum ResolvedResource { pub enum ResolvedResource {
Folder(FolderDto), Folder(FolderDto),
File(FileDto), File(FileDto),
} }
/// Resolves a WebDAV path to a folder or file in a single SQL round-trip. /// Resolves a WebDAV path to a folder or file in a single SQL round-trip.
pub struct PathResolverService { pub struct PathResolverService {
pool: Arc<PgPool>, pool: Arc<PgPool>,
} }
impl PathResolverService { impl PathResolverService {
pub fn new(pool: Arc<PgPool>) -> Self { pub fn new(pool: Arc<PgPool>) -> Self {
Self { pool } Self { pool }
} }
/// Resolve `path` (without leading `/`) to either a folder or a file. /// Resolve `path` (without leading `/`) to either a folder or a file.
/// ///
/// The query uses `UNION ALL … LIMIT 1`: the folder branch is evaluated /// The query uses `UNION ALL … LIMIT 1`: the folder branch is evaluated
/// first, and PG short-circuits if it produces a row. /// first, and PG short-circuits if it produces a row.
pub async fn resolve_path(&self, path: &str) -> Result<ResolvedResource, DomainError> { pub async fn resolve_path(&self, path: &str) -> Result<ResolvedResource, DomainError> {
let path = path.trim_start_matches('/').trim_end_matches('/'); let path = path.trim_start_matches('/').trim_end_matches('/');
if path.is_empty() { if path.is_empty() {
return Err(DomainError::not_found("Resource", "empty path")); return Err(DomainError::not_found("Resource", "empty path"));
} }
// Split into folder_path + filename for the file branch // Split into folder_path + filename for the file branch
let segments: Vec<&str> = path.split('/').filter(|s| !s.is_empty()).collect(); let segments: Vec<&str> = path.split('/').filter(|s| !s.is_empty()).collect();
let filename = segments[segments.len() - 1]; let filename = segments[segments.len() - 1];
let folder_path = if segments.len() > 1 { let folder_path = if segments.len() > 1 {
segments[..segments.len() - 1].join("/") segments[..segments.len() - 1].join("/")
} else { } else {
String::new() String::new()
}; };
// Single round-trip: folder branch ∪ file branch, LIMIT 1. // Single round-trip: folder branch ∪ file branch, LIMIT 1.
// Column order: resource_type, id, name, path, parent_id, user_id, // Column order: resource_type, id, name, path, parent_id, user_id,
// created_at, modified_at, size, mime_type, folder_id // created_at, modified_at, size, mime_type, folder_id
let row = sqlx::query_as::<_, ( let row = sqlx::query_as::<
String, // resource_type _,
String, // id (
String, // name String, // resource_type
String, // path String, // id
Option<String>, // parent_id (folder) / NULL (file) String, // name
Option<String>, // user_id String, // path
i64, // created_at epoch Option<String>, // parent_id (folder) / NULL (file)
i64, // modified_at epoch Option<String>, // user_id
Option<i64>, // size (NULL for folder) i64, // created_at epoch
Option<String>, // mime_type (NULL for folder) i64, // modified_at epoch
Option<String>, // folder_id (NULL for folder) Option<i64>, // size (NULL for folder)
)>( Option<String>, // mime_type (NULL for folder)
r#" Option<String>, // folder_id (NULL for folder)
SELECT resource_type, id, name, path, parent_id, user_id, ),
created_at, modified_at, size, mime_type, folder_id >(
FROM ( r#"
SELECT 'folder'::text AS resource_type, SELECT resource_type, id, name, path, parent_id, user_id,
fo.id::text, created_at, modified_at, size, mime_type, folder_id
fo.name, FROM (
fo.path, SELECT 'folder'::text AS resource_type,
fo.parent_id::text, fo.id::text,
fo.user_id::text, fo.name,
EXTRACT(EPOCH FROM fo.created_at)::bigint AS created_at, fo.path,
EXTRACT(EPOCH FROM fo.updated_at)::bigint AS modified_at, fo.parent_id::text,
NULL::bigint AS size, fo.user_id::text,
NULL::text AS mime_type, EXTRACT(EPOCH FROM fo.created_at)::bigint AS created_at,
NULL::text AS folder_id EXTRACT(EPOCH FROM fo.updated_at)::bigint AS modified_at,
FROM storage.folders fo NULL::bigint AS size,
WHERE fo.path = $1 AND NOT fo.is_trashed NULL::text AS mime_type,
NULL::text AS folder_id
UNION ALL FROM storage.folders fo
WHERE fo.path = $1 AND NOT fo.is_trashed
SELECT 'file'::text AS resource_type,
fi.id::text, UNION ALL
fi.name,
CASE SELECT 'file'::text AS resource_type,
WHEN fo.path IS NOT NULL AND fo.path != '' fi.id::text,
THEN fo.path || '/' || fi.name fi.name,
ELSE fi.name CASE
END AS path, WHEN fo.path IS NOT NULL AND fo.path != ''
NULL::text AS parent_id, THEN fo.path || '/' || fi.name
fi.user_id::text, ELSE fi.name
EXTRACT(EPOCH FROM fi.created_at)::bigint AS created_at, END AS path,
EXTRACT(EPOCH FROM fi.updated_at)::bigint AS modified_at, NULL::text AS parent_id,
fi.size, fi.user_id::text,
fi.mime_type, EXTRACT(EPOCH FROM fi.created_at)::bigint AS created_at,
fi.folder_id::text EXTRACT(EPOCH FROM fi.updated_at)::bigint AS modified_at,
FROM storage.files fi fi.size,
LEFT JOIN storage.folders fo ON fo.id = fi.folder_id fi.mime_type,
WHERE fi.name = $2 fi.folder_id::text
AND ( FROM storage.files fi
($3 = '' AND fi.folder_id IS NULL) LEFT JOIN storage.folders fo ON fo.id = fi.folder_id
OR fo.path = $3 WHERE fi.name = $2
) AND (
AND NOT fi.is_trashed ($3 = '' AND fi.folder_id IS NULL)
) sub OR fo.path = $3
LIMIT 1 )
"#, AND NOT fi.is_trashed
) ) sub
.bind(path) // $1 — full path for folder lookup LIMIT 1
.bind(filename) // $2 — filename for file lookup "#,
.bind(&folder_path) // $3 — parent folder path for file lookup )
.fetch_optional(self.pool.as_ref()) .bind(path) // $1 — full path for folder lookup
.await .bind(filename) // $2 — filename for file lookup
.map_err(|e| DomainError::internal_error("PathResolver", format!("resolve: {e}")))? .bind(&folder_path) // $3 — parent folder path for file lookup
.ok_or_else(|| DomainError::not_found("Resource", path))?; .fetch_optional(self.pool.as_ref())
.await
let (resource_type, id, name, res_path, parent_id, user_id, .map_err(|e| DomainError::internal_error("PathResolver", format!("resolve: {e}")))?
created_at, modified_at, size, mime_type, folder_id) = row; .ok_or_else(|| DomainError::not_found("Resource", path))?;
match resource_type.as_str() { let (
"folder" => Ok(ResolvedResource::Folder(FolderDto { resource_type,
id, id,
name: name.clone(), name,
path: res_path, res_path,
parent_id, parent_id,
owner_id: user_id, user_id,
created_at: created_at as u64, created_at,
modified_at: modified_at as u64, modified_at,
is_root: false, size,
icon_class: "fas fa-folder".to_string(), mime_type,
icon_special_class: "folder-icon".to_string(), folder_id,
category: "Folder".to_string(), ) = row;
})),
_ => { match resource_type.as_str() {
let mime = mime_type.unwrap_or_else(|| "application/octet-stream".to_string()); "folder" => Ok(ResolvedResource::Folder(FolderDto {
let sz = size.unwrap_or(0) as u64; id,
Ok(ResolvedResource::File(FileDto { name: name.clone(),
id, path: res_path,
name: name.clone(), parent_id,
path: res_path, owner_id: user_id,
size: sz, created_at: created_at as u64,
mime_type: mime.clone(), modified_at: modified_at as u64,
folder_id, is_root: false,
created_at: created_at as u64, icon_class: "fas fa-folder".to_string(),
modified_at: modified_at as u64, icon_special_class: "folder-icon".to_string(),
icon_class: icon_class_for(&name, &mime).to_string(), category: "Folder".to_string(),
icon_special_class: icon_special_class_for(&name, &mime).to_string(), })),
category: category_for(&name, &mime).to_string(), _ => {
size_formatted: format_file_size(sz), let mime = mime_type.unwrap_or_else(|| "application/octet-stream".to_string());
owner_id: user_id, let sz = size.unwrap_or(0) as u64;
})) Ok(ResolvedResource::File(FileDto {
} id,
} name: name.clone(),
} path: res_path,
size: sz,
/// Check whether *any* resource (folder or file) exists at the given path. mime_type: mime.clone(),
/// folder_id,
/// Equivalent to `resolve_path(…).is_ok()` but avoids constructing the DTO. created_at: created_at as u64,
pub async fn exists(&self, path: &str) -> Result<bool, DomainError> { modified_at: modified_at as u64,
let path = path.trim_start_matches('/').trim_end_matches('/'); icon_class: icon_class_for(&name, &mime).to_string(),
if path.is_empty() { icon_special_class: icon_special_class_for(&name, &mime).to_string(),
return Ok(false); category: category_for(&name, &mime).to_string(),
} size_formatted: format_file_size(sz),
owner_id: user_id,
let segments: Vec<&str> = path.split('/').filter(|s| !s.is_empty()).collect(); }))
let filename = segments[segments.len() - 1]; }
let folder_path = if segments.len() > 1 { }
segments[..segments.len() - 1].join("/") }
} else {
String::new() /// Check whether *any* resource (folder or file) exists at the given path.
}; ///
/// Equivalent to `resolve_path(…).is_ok()` but avoids constructing the DTO.
let exists = sqlx::query_scalar::<_, bool>( pub async fn exists(&self, path: &str) -> Result<bool, DomainError> {
r#" let path = path.trim_start_matches('/').trim_end_matches('/');
SELECT EXISTS( if path.is_empty() {
SELECT 1 FROM storage.folders return Ok(false);
WHERE path = $1 AND NOT is_trashed }
) OR EXISTS(
SELECT 1 let segments: Vec<&str> = path.split('/').filter(|s| !s.is_empty()).collect();
FROM storage.files fi let filename = segments[segments.len() - 1];
LEFT JOIN storage.folders fo ON fo.id = fi.folder_id let folder_path = if segments.len() > 1 {
WHERE fi.name = $2 segments[..segments.len() - 1].join("/")
AND (($3 = '' AND fi.folder_id IS NULL) OR fo.path = $3) } else {
AND NOT fi.is_trashed String::new()
) };
"#,
) let exists = sqlx::query_scalar::<_, bool>(
.bind(path) r#"
.bind(filename) SELECT EXISTS(
.bind(&folder_path) SELECT 1 FROM storage.folders
.fetch_one(self.pool.as_ref()) WHERE path = $1 AND NOT is_trashed
.await ) OR EXISTS(
.map_err(|e| DomainError::internal_error("PathResolver", format!("exists: {e}")))?; SELECT 1
FROM storage.files fi
Ok(exists) LEFT JOIN storage.folders fo ON fo.id = fi.folder_id
} WHERE fi.name = $2
} AND (($3 = '' AND fi.folder_id IS NULL) OR fo.path = $3)
AND NOT fi.is_trashed
)
"#,
)
.bind(path)
.bind(filename)
.bind(&folder_path)
.fetch_one(self.pool.as_ref())
.await
.map_err(|e| DomainError::internal_error("PathResolver", format!("exists: {e}")))?;
Ok(exists)
}
}
+140 -148
View File
@@ -1,148 +1,140 @@
//! HttpOnly cookie helpers for secure token transport. //! HttpOnly cookie helpers for secure token transport.
//! //!
//! Tokens are set as `HttpOnly; SameSite=Lax` cookies so that //! Tokens are set as `HttpOnly; SameSite=Lax` cookies so that
//! browser-based JavaScript cannot read them (mitigates XSS token theft). //! browser-based JavaScript cannot read them (mitigates XSS token theft).
//! The `Secure` flag is controlled by the `OXICLOUD_COOKIE_SECURE` env var //! The `Secure` flag is controlled by the `OXICLOUD_COOKIE_SECURE` env var
//! (default: auto-detect from `OXICLOUD_BASE_URL`). //! (default: auto-detect from `OXICLOUD_BASE_URL`).
//! //!
//! A companion **non-HttpOnly** CSRF cookie (`oxicloud_csrf`) is set //! A companion **non-HttpOnly** CSRF cookie (`oxicloud_csrf`) is set
//! alongside the auth cookies. The frontend must read it and echo its //! alongside the auth cookies. The frontend must read it and echo its
//! value back as `X-CSRF-Token` on every state-changing request. //! value back as `X-CSRF-Token` on every state-changing request.
//! A middleware (`csrf_middleware`) validates the match. //! A middleware (`csrf_middleware`) validates the match.
//! //!
//! DAV clients continue to use `Authorization: Basic` with app passwords //! DAV clients continue to use `Authorization: Basic` with app passwords
//! and are completely unaffected by this mechanism. //! and are completely unaffected by this mechanism.
use axum::http::header::SET_COOKIE; use axum::http::header::SET_COOKIE;
use axum::http::{HeaderMap, HeaderValue}; use axum::http::{HeaderMap, HeaderValue};
/// Cookie name for the JWT access token. /// Cookie name for the JWT access token.
pub const ACCESS_COOKIE: &str = "oxicloud_access"; pub const ACCESS_COOKIE: &str = "oxicloud_access";
/// Cookie name for the opaque refresh token. /// Cookie name for the opaque refresh token.
pub const REFRESH_COOKIE: &str = "oxicloud_refresh"; pub const REFRESH_COOKIE: &str = "oxicloud_refresh";
/// Cookie name for the CSRF double-submit token (readable by JS). /// Cookie name for the CSRF double-submit token (readable by JS).
pub const CSRF_COOKIE: &str = "oxicloud_csrf"; pub const CSRF_COOKIE: &str = "oxicloud_csrf";
/// Header the frontend must send with the CSRF token value. /// Header the frontend must send with the CSRF token value.
pub const CSRF_HEADER: &str = "x-csrf-token"; pub const CSRF_HEADER: &str = "x-csrf-token";
/// Whether the `Secure` flag should be set on cookies. /// Whether the `Secure` flag should be set on cookies.
/// Auto-detected from `OXICLOUD_BASE_URL` (if it starts with `https`) /// Auto-detected from `OXICLOUD_BASE_URL` (if it starts with `https`)
/// or overridden with `OXICLOUD_COOKIE_SECURE=true|false`. /// or overridden with `OXICLOUD_COOKIE_SECURE=true|false`.
fn cookie_secure() -> bool { fn cookie_secure() -> bool {
if let Ok(v) = std::env::var("OXICLOUD_COOKIE_SECURE") { if let Ok(v) = std::env::var("OXICLOUD_COOKIE_SECURE") {
return v == "true" || v == "1"; return v == "true" || v == "1";
} }
// Auto-detect from base URL // Auto-detect from base URL
std::env::var("OXICLOUD_BASE_URL") std::env::var("OXICLOUD_BASE_URL")
.map(|u| u.starts_with("https")) .map(|u| u.starts_with("https"))
.unwrap_or(false) .unwrap_or(false)
} }
/// Build a `Set-Cookie` header value. /// Build a `Set-Cookie` header value.
fn build_cookie(name: &str, value: &str, path: &str, max_age_secs: i64) -> String { fn build_cookie(name: &str, value: &str, path: &str, max_age_secs: i64) -> String {
let secure = if cookie_secure() { "; Secure" } else { "" }; let secure = if cookie_secure() { "; Secure" } else { "" };
format!( format!("{name}={value}; HttpOnly; SameSite=Lax; Path={path}; Max-Age={max_age_secs}{secure}",)
"{name}={value}; HttpOnly; SameSite=Lax; Path={path}; Max-Age={max_age_secs}{secure}", }
)
} /// Append `Set-Cookie` headers for both access and refresh tokens.
///
/// Append `Set-Cookie` headers for both access and refresh tokens. /// The access cookie covers all paths (`/`) because the API lives under
/// /// `/api`, CalDAV under `/caldav`, WebDAV under `/webdav`, etc.
/// The access cookie covers all paths (`/`) because the API lives under ///
/// `/api`, CalDAV under `/caldav`, WebDAV under `/webdav`, etc. /// The refresh cookie is restricted to `/api/auth` so it is only sent
/// /// when the client explicitly calls the refresh or logout endpoints.
/// The refresh cookie is restricted to `/api/auth` so it is only sent pub fn append_auth_cookies(
/// when the client explicitly calls the refresh or logout endpoints. headers: &mut HeaderMap,
pub fn append_auth_cookies( access_token: &str,
headers: &mut HeaderMap, refresh_token: &str,
access_token: &str, access_expiry_secs: i64,
refresh_token: &str, refresh_expiry_secs: i64,
access_expiry_secs: i64, ) {
refresh_expiry_secs: i64, if let Ok(val) = HeaderValue::from_str(&build_cookie(
) { ACCESS_COOKIE,
if let Ok(val) = HeaderValue::from_str(&build_cookie( access_token,
ACCESS_COOKIE, "/",
access_token, access_expiry_secs,
"/", )) {
access_expiry_secs, headers.append(SET_COOKIE, val);
)) { }
headers.append(SET_COOKIE, val); if let Ok(val) = HeaderValue::from_str(&build_cookie(
} REFRESH_COOKIE,
if let Ok(val) = HeaderValue::from_str(&build_cookie( refresh_token,
REFRESH_COOKIE, "/api/auth",
refresh_token, refresh_expiry_secs,
"/api/auth", )) {
refresh_expiry_secs, headers.append(SET_COOKIE, val);
)) { }
headers.append(SET_COOKIE, val); }
}
} /// Append `Set-Cookie` headers that immediately expire both auth cookies,
/// effectively logging the user out on the browser side.
/// Append `Set-Cookie` headers that immediately expire both auth cookies, pub fn append_clear_cookies(headers: &mut HeaderMap) {
/// effectively logging the user out on the browser side. for (name, path) in [(ACCESS_COOKIE, "/"), (REFRESH_COOKIE, "/api/auth")] {
pub fn append_clear_cookies(headers: &mut HeaderMap) { let secure = if cookie_secure() { "; Secure" } else { "" };
for (name, path) in [(ACCESS_COOKIE, "/"), (REFRESH_COOKIE, "/api/auth")] { let val = format!("{name}=; HttpOnly; SameSite=Lax; Path={path}; Max-Age=0{secure}",);
let secure = if cookie_secure() { "; Secure" } else { "" }; if let Ok(hv) = HeaderValue::from_str(&val) {
let val = format!( headers.append(SET_COOKIE, hv);
"{name}=; HttpOnly; SameSite=Lax; Path={path}; Max-Age=0{secure}", }
); }
if let Ok(hv) = HeaderValue::from_str(&val) { }
headers.append(SET_COOKIE, hv);
} /// Extract a named cookie value from the `Cookie` request header.
} pub fn extract_cookie_value(headers: &HeaderMap, name: &str) -> Option<String> {
} let cookie_header = headers.get(axum::http::header::COOKIE)?;
let cookie_str = cookie_header.to_str().ok()?;
/// Extract a named cookie value from the `Cookie` request header.
pub fn extract_cookie_value(headers: &HeaderMap, name: &str) -> Option<String> { for pair in cookie_str.split(';') {
let cookie_header = headers.get(axum::http::header::COOKIE)?; let pair = pair.trim();
let cookie_str = cookie_header.to_str().ok()?; if let Some(val) = pair.strip_prefix(name) {
let val = val.strip_prefix('=')?;
for pair in cookie_str.split(';') { if !val.is_empty() {
let pair = pair.trim(); return Some(val.to_string());
if let Some(val) = pair.strip_prefix(name) { }
let val = val.strip_prefix('=')?; }
if !val.is_empty() { }
return Some(val.to_string()); None
} }
}
} // ────────────────────────────────────────────────────────────
None // CSRF double-submit cookie helpers
} // ────────────────────────────────────────────────────────────
// ──────────────────────────────────────────────────────────── /// Generate a cryptographically random CSRF token (128-bit UUIDv4, hex-like).
// CSRF double-submit cookie helpers pub fn generate_csrf_token() -> String {
// ──────────────────────────────────────────────────────────── uuid::Uuid::new_v4().to_string()
}
/// Generate a cryptographically random CSRF token (128-bit UUIDv4, hex-like).
pub fn generate_csrf_token() -> String { /// Build a **non-HttpOnly** CSRF cookie so that frontend JS can read it
uuid::Uuid::new_v4().to_string() /// via `document.cookie` and echo it back in the `X-CSRF-Token` header.
} fn build_csrf_cookie(value: &str, max_age_secs: i64) -> String {
let secure = if cookie_secure() { "; Secure" } else { "" };
/// Build a **non-HttpOnly** CSRF cookie so that frontend JS can read it format!("{CSRF_COOKIE}={value}; SameSite=Lax; Path=/; Max-Age={max_age_secs}{secure}",)
/// via `document.cookie` and echo it back in the `X-CSRF-Token` header. }
fn build_csrf_cookie(value: &str, max_age_secs: i64) -> String {
let secure = if cookie_secure() { "; Secure" } else { "" }; /// Append a CSRF double-submit cookie alongside the auth cookies.
format!( /// Should be called in every endpoint that also sets auth cookies.
"{CSRF_COOKIE}={value}; SameSite=Lax; Path=/; Max-Age={max_age_secs}{secure}", pub fn append_csrf_cookie(headers: &mut HeaderMap, access_expiry_secs: i64) {
) let token = generate_csrf_token();
} if let Ok(val) = HeaderValue::from_str(&build_csrf_cookie(&token, access_expiry_secs)) {
headers.append(SET_COOKIE, val);
/// Append a CSRF double-submit cookie alongside the auth cookies. }
/// Should be called in every endpoint that also sets auth cookies. }
pub fn append_csrf_cookie(headers: &mut HeaderMap, access_expiry_secs: i64) {
let token = generate_csrf_token(); /// Clear the CSRF cookie (on logout).
if let Ok(val) = HeaderValue::from_str(&build_csrf_cookie(&token, access_expiry_secs)) { pub fn append_clear_csrf_cookie(headers: &mut HeaderMap) {
headers.append(SET_COOKIE, val); let secure = if cookie_secure() { "; Secure" } else { "" };
} let val = format!("{CSRF_COOKIE}=; SameSite=Lax; Path=/; Max-Age=0{secure}",);
} if let Ok(hv) = HeaderValue::from_str(&val) {
headers.append(SET_COOKIE, hv);
/// Clear the CSRF cookie (on logout). }
pub fn append_clear_csrf_cookie(headers: &mut HeaderMap) { }
let secure = if cookie_secure() { "; Secure" } else { "" };
let val = format!(
"{CSRF_COOKIE}=; SameSite=Lax; Path=/; Max-Age=0{secure}",
);
if let Ok(hv) = HeaderValue::from_str(&val) {
headers.append(SET_COOKIE, hv);
}
}
+3 -12
View File
@@ -197,10 +197,7 @@ async fn login(
auth_response.expires_in, auth_response.expires_in,
state.core.config.auth.refresh_token_expiry_secs, state.core.config.auth.refresh_token_expiry_secs,
); );
cookie_auth::append_csrf_cookie( cookie_auth::append_csrf_cookie(response.headers_mut(), auth_response.expires_in);
response.headers_mut(),
auth_response.expires_in,
);
Ok(response) Ok(response)
} }
Err(err) => { Err(err) => {
@@ -253,10 +250,7 @@ async fn refresh_token(
auth_response.expires_in, auth_response.expires_in,
state.core.config.auth.refresh_token_expiry_secs, state.core.config.auth.refresh_token_expiry_secs,
); );
cookie_auth::append_csrf_cookie( cookie_auth::append_csrf_cookie(response.headers_mut(), auth_response.expires_in);
response.headers_mut(),
auth_response.expires_in,
);
Ok(response) Ok(response)
} }
@@ -522,9 +516,6 @@ async fn oidc_exchange(
auth_response.expires_in, auth_response.expires_in,
state.core.config.auth.refresh_token_expiry_secs, state.core.config.auth.refresh_token_expiry_secs,
); );
cookie_auth::append_csrf_cookie( cookie_auth::append_csrf_cookie(response.headers_mut(), auth_response.expires_in);
response.headers_mut(),
auth_response.expires_in,
);
Ok(response) Ok(response)
} }
+28 -50
View File
@@ -52,11 +52,10 @@ pub fn caldav_routes() -> Router<Arc<AppState>> {
/// Creates RFC 6764 well-known discovery routes. /// Creates RFC 6764 well-known discovery routes.
/// These are public (no auth) and simply redirect to the CalDAV root. /// These are public (no auth) and simply redirect to the CalDAV root.
pub fn well_known_routes() -> Router<Arc<AppState>> { pub fn well_known_routes() -> Router<Arc<AppState>> {
Router::new() Router::new().route(
.route( "/.well-known/caldav",
"/.well-known/caldav", axum::routing::any(handle_well_known_caldav),
axum::routing::any(handle_well_known_caldav), )
)
} }
async fn handle_well_known_caldav() -> Response<Body> { async fn handle_well_known_caldav() -> Response<Body> {
@@ -114,13 +113,9 @@ fn extract_caldav_path(uri_path: &str) -> String {
} else if uri_path.ends_with("/caldav") { } else if uri_path.ends_with("/caldav") {
"" ""
} else { } else {
uri_path uri_path.trim_start_matches('/').trim_end_matches('/')
.trim_start_matches('/')
.trim_end_matches('/')
}; };
percent_decode_str(encoded) percent_decode_str(encoded).decode_utf8_lossy().into_owned()
.decode_utf8_lossy()
.into_owned()
} }
// ─── Helper: extract user from request ─────────────────────────────── // ─── Helper: extract user from request ───────────────────────────────
@@ -198,9 +193,7 @@ async fn handle_propfind(
calendar_service calendar_service
.list_my_calendars(&user.id) .list_my_calendars(&user.id)
.await .await
.map_err(|e| { .map_err(|e| AppError::internal_error(format!("Failed to list calendars: {}", e)))?
AppError::internal_error(format!("Failed to list calendars: {}", e))
})?
}; };
let base_href = "/caldav/"; let base_href = "/caldav/";
@@ -221,9 +214,7 @@ async fn handle_propfind(
.unwrap()) .unwrap())
} else if path.starts_with("principals/") || path == "principals" { } else if path.starts_with("principals/") || path == "principals" {
// Principal resource — return user principal properties // Principal resource — return user principal properties
let username = path let username = path.strip_prefix("principals/").unwrap_or(&user.username);
.strip_prefix("principals/")
.unwrap_or(&user.username);
let username = if username.is_empty() { let username = if username.is_empty() {
&user.username &user.username
} else { } else {
@@ -255,9 +246,7 @@ async fn handle_propfind(
if parts.len() == 1 { if parts.len() == 1 {
// Single path segment: try as calendar ID first, fall back to user home // Single path segment: try as calendar ID first, fall back to user home
let calendar_result = calendar_service let calendar_result = calendar_service.get_calendar(first_segment, &user.id).await;
.get_calendar(first_segment, &user.id)
.await;
if let Ok(calendar) = calendar_result { if let Ok(calendar) = calendar_result {
// Valid calendar ID — return calendar collection // Valid calendar ID — return calendar collection
@@ -281,9 +270,7 @@ async fn handle_propfind(
base_href, base_href,
&depth, &depth,
) )
.map_err(|e| { .map_err(|e| AppError::internal_error(format!("Failed to generate XML: {}", e)))?;
AppError::internal_error(format!("Failed to generate XML: {}", e))
})?;
Ok(Response::builder() Ok(Response::builder()
.status(StatusCode::MULTI_STATUS) .status(StatusCode::MULTI_STATUS)
@@ -293,12 +280,13 @@ async fn handle_propfind(
} else { } else {
// Not a calendar ID — treat as user calendar home (e.g. /caldav/{username}/) // Not a calendar ID — treat as user calendar home (e.g. /caldav/{username}/)
// List all calendars for this user // List all calendars for this user
let calendars = calendar_service let calendars =
.list_my_calendars(&user.id) calendar_service
.await .list_my_calendars(&user.id)
.map_err(|e| { .await
AppError::internal_error(format!("Failed to list calendars: {}", e)) .map_err(|e| {
})?; AppError::internal_error(format!("Failed to list calendars: {}", e))
})?;
let base_href = &format!("/caldav/{}/", first_segment); let base_href = &format!("/caldav/{}/", first_segment);
let mut response_body = Vec::new(); let mut response_body = Vec::new();
@@ -309,9 +297,7 @@ async fn handle_propfind(
&propfind_request, &propfind_request,
base_href, base_href,
) )
.map_err(|e| { .map_err(|e| AppError::internal_error(format!("Failed to generate XML: {}", e)))?;
AppError::internal_error(format!("Failed to generate XML: {}", e))
})?;
Ok(Response::builder() Ok(Response::builder()
.status(StatusCode::MULTI_STATUS) .status(StatusCode::MULTI_STATUS)
@@ -324,9 +310,7 @@ async fn handle_propfind(
let rest = parts[1]; let rest = parts[1];
// Check if first_segment is a valid calendar ID // Check if first_segment is a valid calendar ID
let calendar_result = calendar_service let calendar_result = calendar_service.get_calendar(first_segment, &user.id).await;
.get_calendar(first_segment, &user.id)
.await;
let (calendar_id, event_path) = if calendar_result.is_ok() { let (calendar_id, event_path) = if calendar_result.is_ok() {
// first_segment is a calendar ID, rest is event path // first_segment is a calendar ID, rest is event path
@@ -341,9 +325,7 @@ async fn handle_propfind(
let cal = calendar_service let cal = calendar_service
.get_calendar(sub_parts[0], &user.id) .get_calendar(sub_parts[0], &user.id)
.await .await
.map_err(|e| { .map_err(|e| AppError::not_found(format!("Calendar not found: {}", e)))?;
AppError::not_found(format!("Calendar not found: {}", e))
})?;
let events = if depth != "0" { let events = if depth != "0" {
calendar_service calendar_service
@@ -354,8 +336,7 @@ async fn handle_propfind(
vec![] vec![]
}; };
let base_href = let base_href = &format!("/caldav/{}/{}/", first_segment, sub_parts[0]);
&format!("/caldav/{}/{}/", first_segment, sub_parts[0]);
let mut response_body = Vec::new(); let mut response_body = Vec::new();
CalDavAdapter::generate_calendar_collection_propfind( CalDavAdapter::generate_calendar_collection_propfind(
@@ -387,16 +368,12 @@ async fn handle_propfind(
let events = calendar_service let events = calendar_service
.list_events(calendar_id, None, None, &user.id) .list_events(calendar_id, None, None, &user.id)
.await .await
.map_err(|e| { .map_err(|e| AppError::internal_error(format!("Failed to list events: {}", e)))?;
AppError::internal_error(format!("Failed to list events: {}", e))
})?;
let event = events let event = events
.iter() .iter()
.find(|e| e.ical_uid == ical_uid) .find(|e| e.ical_uid == ical_uid)
.ok_or_else(|| { .ok_or_else(|| AppError::not_found(format!("Event not found: {}", ical_uid)))?;
AppError::not_found(format!("Event not found: {}", ical_uid))
})?;
let base_href = &format!("/caldav/{}/", calendar_id); let base_href = &format!("/caldav/{}/", calendar_id);
let report_type = CalDavReportType::CalendarMultiget { let report_type = CalDavReportType::CalendarMultiget {
@@ -411,9 +388,7 @@ async fn handle_propfind(
&report_type, &report_type,
base_href, base_href,
) )
.map_err(|e| { .map_err(|e| AppError::internal_error(format!("Failed to generate XML: {}", e)))?;
AppError::internal_error(format!("Failed to generate XML: {}", e))
})?;
Ok(Response::builder() Ok(Response::builder()
.status(StatusCode::MULTI_STATUS) .status(StatusCode::MULTI_STATUS)
@@ -718,7 +693,10 @@ fn generate_event_ical(event: &crate::application::dtos::calendar_dto::CalendarE
} }
/// Writes a VEVENT block directly into `buf` — zero intermediate allocations. /// Writes a VEVENT block directly into `buf` — zero intermediate allocations.
fn write_vevent(buf: &mut String, event: &crate::application::dtos::calendar_dto::CalendarEventDto) { fn write_vevent(
buf: &mut String,
event: &crate::application::dtos::calendar_dto::CalendarEventDto,
) {
let _ = write!( let _ = write!(
buf, buf,
"BEGIN:VEVENT\r\nUID:{}\r\nSUMMARY:{}\r\nDTSTART:{}\r\nDTEND:{}\r\n", "BEGIN:VEVENT\r\nUID:{}\r\nSUMMARY:{}\r\nDTSTART:{}\r\nDTEND:{}\r\n",
@@ -99,9 +99,7 @@ fn extract_carddav_path(uri_path: &str) -> String {
} else if uri_path.ends_with("/carddav") { } else if uri_path.ends_with("/carddav") {
"" ""
} else { } else {
uri_path uri_path.trim_start_matches('/').trim_end_matches('/')
.trim_start_matches('/')
.trim_end_matches('/')
}; };
percent_encoding::percent_decode_str(encoded) percent_encoding::percent_decode_str(encoded)
.decode_utf8_lossy() .decode_utf8_lossy()
+225 -227
View File
@@ -1,227 +1,225 @@
//! HTTP handlers for OAuth 2.0 Device Authorization Grant (RFC 8628). //! HTTP handlers for OAuth 2.0 Device Authorization Grant (RFC 8628).
//! //!
//! Endpoints: //! Endpoints:
//! POST /api/auth/device/authorize — Client starts the device flow (public) //! POST /api/auth/device/authorize — Client starts the device flow (public)
//! GET /api/auth/device/verify — Check user_code validity (authenticated) //! GET /api/auth/device/verify — Check user_code validity (authenticated)
//! POST /api/auth/device/verify — User approves/denies (authenticated) //! POST /api/auth/device/verify — User approves/denies (authenticated)
//! POST /api/auth/device/token — Client polls for tokens (public) //! POST /api/auth/device/token — Client polls for tokens (public)
//! GET /api/auth/device/devices — List user's authorized devices (authenticated) //! GET /api/auth/device/devices — List user's authorized devices (authenticated)
//! DELETE /api/auth/device/devices/{id} — Revoke a device (authenticated) //! DELETE /api/auth/device/devices/{id} — Revoke a device (authenticated)
use axum::{ use axum::{
Router, Router,
extract::{Json, Path, Query, State}, extract::{Json, Path, Query, State},
http::StatusCode, http::StatusCode,
response::IntoResponse, response::IntoResponse,
routing::{delete, get, post}, routing::{delete, get, post},
}; };
use std::sync::Arc; use std::sync::Arc;
use crate::application::dtos::device_auth_dto::*; use crate::application::dtos::device_auth_dto::*;
use crate::application::services::device_auth_service::DeviceAuthService; use crate::application::services::device_auth_service::DeviceAuthService;
use crate::common::di::AppState; use crate::common::di::AppState;
use crate::interfaces::errors::AppError; use crate::interfaces::errors::AppError;
use crate::interfaces::middleware::auth::AuthUser; use crate::interfaces::middleware::auth::AuthUser;
/// Create the device auth router. /// Create the device auth router.
/// ///
/// Public endpoints (no auth middleware): authorize, token /// Public endpoints (no auth middleware): authorize, token
/// Protected endpoints (behind auth middleware): verify (GET+POST), devices /// Protected endpoints (behind auth middleware): verify (GET+POST), devices
pub fn device_auth_public_routes() -> Router<Arc<AppState>> { pub fn device_auth_public_routes() -> Router<Arc<AppState>> {
Router::new() Router::new()
// Client-facing endpoints (no auth needed — the client doesn't have tokens yet) // Client-facing endpoints (no auth needed — the client doesn't have tokens yet)
.route("/authorize", post(device_authorize)) .route("/authorize", post(device_authorize))
.route("/token", post(device_token)) .route("/token", post(device_token))
} }
pub fn device_auth_protected_routes() -> Router<Arc<AppState>> { pub fn device_auth_protected_routes() -> Router<Arc<AppState>> {
Router::new() Router::new()
// User-facing endpoints (require valid session) // User-facing endpoints (require valid session)
.route("/verify", get(device_verify_info)) .route("/verify", get(device_verify_info))
.route("/verify", post(device_verify_action)) .route("/verify", post(device_verify_action))
.route("/devices", get(list_devices)) .route("/devices", get(list_devices))
.route("/devices/{id}", delete(revoke_device)) .route("/devices/{id}", delete(revoke_device))
} }
// ============================================================================ // ============================================================================
// POST /api/auth/device/authorize — Client initiates the device flow // POST /api/auth/device/authorize — Client initiates the device flow
// ============================================================================ // ============================================================================
/// Client sends: `{ "client_name": "rclone", "scope": "webdav" }` /// Client sends: `{ "client_name": "rclone", "scope": "webdav" }`
/// Server returns: device_code, user_code, verification_uri, etc. /// Server returns: device_code, user_code, verification_uri, etc.
async fn device_authorize( async fn device_authorize(
State(state): State<Arc<AppState>>, State(state): State<Arc<AppState>>,
Json(body): Json<DeviceAuthorizeRequestDto>, Json(body): Json<DeviceAuthorizeRequestDto>,
) -> Result<impl IntoResponse, AppError> { ) -> Result<impl IntoResponse, AppError> {
let device_service = get_device_service(&state)?; let device_service = get_device_service(&state)?;
let response = device_service.initiate(body).await.map_err(|e| { let response = device_service.initiate(body).await.map_err(|e| {
tracing::error!("Device authorize failed: {}", e); tracing::error!("Device authorize failed: {}", e);
AppError::from(e) AppError::from(e)
})?; })?;
Ok((StatusCode::OK, Json(response))) Ok((StatusCode::OK, Json(response)))
} }
// ============================================================================ // ============================================================================
// POST /api/auth/device/token — Client polls for tokens // POST /api/auth/device/token — Client polls for tokens
// ============================================================================ // ============================================================================
/// Client sends: `{ "device_code": "...", "grant_type": "urn:ietf:params:oauth:grant-type:device_code" }` /// Client sends: `{ "device_code": "...", "grant_type": "urn:ietf:params:oauth:grant-type:device_code" }`
/// Returns tokens on success, or RFC 8628 error codes while pending. /// Returns tokens on success, or RFC 8628 error codes while pending.
async fn device_token( async fn device_token(
State(state): State<Arc<AppState>>, State(state): State<Arc<AppState>>,
Json(body): Json<DeviceTokenRequestDto>, Json(body): Json<DeviceTokenRequestDto>,
) -> Result<impl IntoResponse, impl IntoResponse> { ) -> Result<impl IntoResponse, impl IntoResponse> {
let device_service = match get_device_service(&state) { let device_service = match get_device_service(&state) {
Ok(svc) => svc, Ok(svc) => svc,
Err(e) => return Err(e.into_response()), Err(e) => return Err(e.into_response()),
}; };
// Validate grant_type if provided (RFC compliance) // Validate grant_type if provided (RFC compliance)
if !body.grant_type.is_empty() if !body.grant_type.is_empty()
&& body.grant_type != "urn:ietf:params:oauth:grant-type:device_code" && body.grant_type != "urn:ietf:params:oauth:grant-type:device_code"
{ {
let error_body = serde_json::json!({ let error_body = serde_json::json!({
"error": "unsupported_grant_type", "error": "unsupported_grant_type",
"error_description": "grant_type must be urn:ietf:params:oauth:grant-type:device_code" "error_description": "grant_type must be urn:ietf:params:oauth:grant-type:device_code"
}); });
return Err((StatusCode::BAD_REQUEST, Json(error_body)).into_response()); return Err((StatusCode::BAD_REQUEST, Json(error_body)).into_response());
} }
match device_service.poll(&body.device_code).await { match device_service.poll(&body.device_code).await {
Ok(tokens) => Ok((StatusCode::OK, Json(tokens)).into_response()), Ok(tokens) => Ok((StatusCode::OK, Json(tokens)).into_response()),
Err(poll_err) => { Err(poll_err) => {
let status = StatusCode::from_u16(poll_err.http_status()) let status =
.unwrap_or(StatusCode::BAD_REQUEST); StatusCode::from_u16(poll_err.http_status()).unwrap_or(StatusCode::BAD_REQUEST);
let error_body = serde_json::json!({ let error_body = serde_json::json!({
"error": poll_err.error_code(), "error": poll_err.error_code(),
"error_description": poll_err.description() "error_description": poll_err.description()
}); });
Err((status, Json(error_body)).into_response()) Err((status, Json(error_body)).into_response())
} }
} }
} }
// ============================================================================ // ============================================================================
// GET /api/auth/device/verify?code=ABCD-1234 — Check if user_code is valid // GET /api/auth/device/verify?code=ABCD-1234 — Check if user_code is valid
// ============================================================================ // ============================================================================
#[derive(serde::Deserialize)] #[derive(serde::Deserialize)]
pub struct VerifyQuery { pub struct VerifyQuery {
#[serde(default)] #[serde(default)]
pub code: String, pub code: String,
} }
async fn device_verify_info( async fn device_verify_info(
State(state): State<Arc<AppState>>, State(state): State<Arc<AppState>>,
_auth_user: AuthUser, _auth_user: AuthUser,
Query(query): Query<VerifyQuery>, Query(query): Query<VerifyQuery>,
) -> Result<impl IntoResponse, AppError> { ) -> Result<impl IntoResponse, AppError> {
let device_service = get_device_service(&state)?; let device_service = get_device_service(&state)?;
let info = device_service let info = device_service
.verify_user_code(&query.code) .verify_user_code(&query.code)
.await .await
.map_err(|e| { .map_err(|e| {
tracing::warn!("Device verify lookup failed: {}", e); tracing::warn!("Device verify lookup failed: {}", e);
AppError::from(e) AppError::from(e)
})?; })?;
Ok((StatusCode::OK, Json(info))) Ok((StatusCode::OK, Json(info)))
} }
// ============================================================================ // ============================================================================
// POST /api/auth/device/verify — User approves or denies // POST /api/auth/device/verify — User approves or denies
// ============================================================================ // ============================================================================
async fn device_verify_action( async fn device_verify_action(
State(state): State<Arc<AppState>>, State(state): State<Arc<AppState>>,
auth_user: AuthUser, auth_user: AuthUser,
Json(body): Json<DeviceVerifyRequestDto>, Json(body): Json<DeviceVerifyRequestDto>,
) -> Result<impl IntoResponse, AppError> { ) -> Result<impl IntoResponse, AppError> {
let device_service = get_device_service(&state)?; let device_service = get_device_service(&state)?;
match body.action.to_lowercase().as_str() { match body.action.to_lowercase().as_str() {
"approve" | "allow" | "accept" => { "approve" | "allow" | "accept" => {
device_service device_service
.approve(&body.user_code, &auth_user.id) .approve(&body.user_code, &auth_user.id)
.await .await
.map_err(|e| { .map_err(|e| {
tracing::error!("Device approve failed: {}", e); tracing::error!("Device approve failed: {}", e);
AppError::from(e) AppError::from(e)
})?; })?;
Ok(( Ok((
StatusCode::OK, StatusCode::OK,
Json(serde_json::json!({ "status": "approved" })), Json(serde_json::json!({ "status": "approved" })),
)) ))
} }
"deny" | "reject" | "cancel" => { "deny" | "reject" | "cancel" => {
device_service.deny(&body.user_code).await.map_err(|e| { device_service.deny(&body.user_code).await.map_err(|e| {
tracing::error!("Device deny failed: {}", e); tracing::error!("Device deny failed: {}", e);
AppError::from(e) AppError::from(e)
})?; })?;
Ok(( Ok((
StatusCode::OK, StatusCode::OK,
Json(serde_json::json!({ "status": "denied" })), Json(serde_json::json!({ "status": "denied" })),
)) ))
} }
_ => Err(AppError::bad_request( _ => Err(AppError::bad_request("action must be 'approve' or 'deny'")),
"action must be 'approve' or 'deny'", }
)), }
}
} // ============================================================================
// GET /api/auth/device/devices — List user's authorized devices
// ============================================================================ // ============================================================================
// GET /api/auth/device/devices — List user's authorized devices
// ============================================================================ async fn list_devices(
State(state): State<Arc<AppState>>,
async fn list_devices( auth_user: AuthUser,
State(state): State<Arc<AppState>>, ) -> Result<impl IntoResponse, AppError> {
auth_user: AuthUser, let device_service = get_device_service(&state)?;
) -> Result<impl IntoResponse, AppError> {
let device_service = get_device_service(&state)?; let devices = device_service
.list_user_devices(&auth_user.id)
let devices = device_service .await
.list_user_devices(&auth_user.id) .map_err(|e| {
.await tracing::error!("List devices failed: {}", e);
.map_err(|e| { AppError::from(e)
tracing::error!("List devices failed: {}", e); })?;
AppError::from(e)
})?; Ok((StatusCode::OK, Json(devices)))
}
Ok((StatusCode::OK, Json(devices)))
} // ============================================================================
// DELETE /api/auth/device/devices/{id} — Revoke a device authorization
// ============================================================================ // ============================================================================
// DELETE /api/auth/device/devices/{id} — Revoke a device authorization
// ============================================================================ async fn revoke_device(
State(state): State<Arc<AppState>>,
async fn revoke_device( auth_user: AuthUser,
State(state): State<Arc<AppState>>, Path(device_id): Path<String>,
auth_user: AuthUser, ) -> Result<impl IntoResponse, AppError> {
Path(device_id): Path<String>, let device_service = get_device_service(&state)?;
) -> Result<impl IntoResponse, AppError> {
let device_service = get_device_service(&state)?; device_service
.revoke_device(&device_id, &auth_user.id)
device_service .await
.revoke_device(&device_id, &auth_user.id) .map_err(|e| {
.await tracing::error!("Revoke device failed: {}", e);
.map_err(|e| { AppError::from(e)
tracing::error!("Revoke device failed: {}", e); })?;
AppError::from(e)
})?; Ok(StatusCode::NO_CONTENT)
}
Ok(StatusCode::NO_CONTENT)
} // ============================================================================
// Helper
// ============================================================================ // ============================================================================
// Helper
// ============================================================================ fn get_device_service(state: &AppState) -> Result<&Arc<DeviceAuthService>, AppError> {
state
fn get_device_service(state: &AppState) -> Result<&Arc<DeviceAuthService>, AppError> { .device_auth_service
state .as_ref()
.device_auth_service .ok_or_else(|| AppError::internal_error("Device authorization service not configured"))
.as_ref() }
.ok_or_else(|| AppError::internal_error("Device authorization service not configured"))
}
+1 -1
View File
@@ -2,11 +2,11 @@ pub mod admin_handler;
pub mod app_password_handler; pub mod app_password_handler;
pub mod auth_handler; pub mod auth_handler;
pub mod batch_handler; pub mod batch_handler;
pub mod device_auth_handler;
pub mod caldav_handler; pub mod caldav_handler;
pub mod carddav_handler; pub mod carddav_handler;
pub mod chunked_upload_handler; pub mod chunked_upload_handler;
pub mod dedup_handler; pub mod dedup_handler;
pub mod device_auth_handler;
pub mod favorites_handler; pub mod favorites_handler;
pub mod file_handler; pub mod file_handler;
pub mod folder_handler; pub mod folder_handler;
+82 -26
View File
@@ -28,7 +28,7 @@ use crate::common::di::AppState;
use crate::infrastructure::services::path_resolver_service::ResolvedResource; use crate::infrastructure::services::path_resolver_service::ResolvedResource;
use crate::interfaces::errors::AppError; use crate::interfaces::errors::AppError;
use crate::interfaces::middleware::auth::CurrentUser; use crate::interfaces::middleware::auth::CurrentUser;
use percent_encoding::{percent_decode_str, utf8_percent_encode, NON_ALPHANUMERIC, AsciiSet}; use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, percent_decode_str, utf8_percent_encode};
use std::sync::Arc; use std::sync::Arc;
/// Characters that MUST NOT be percent-encoded inside a URI path segment. /// Characters that MUST NOT be percent-encoded inside a URI path segment.
@@ -115,9 +115,7 @@ fn extract_webdav_path(uri: &axum::http::Uri) -> String {
trimmed.trim_end_matches('/') trimmed.trim_end_matches('/')
}; };
// Decode percent-encoded characters (e.g. %20 → space) // Decode percent-encoded characters (e.g. %20 → space)
percent_decode_str(encoded) percent_decode_str(encoded).decode_utf8_lossy().into_owned()
.decode_utf8_lossy()
.into_owned()
} }
async fn handle_webdav_methods_root( async fn handle_webdav_methods_root(
@@ -324,8 +322,13 @@ async fn handle_propfind(
let mut xml_writer = Writer::new(&mut buf); let mut xml_writer = Writer::new(&mut buf);
WebDavAdapter::write_multistatus_start(&mut xml_writer) WebDavAdapter::write_multistatus_start(&mut xml_writer)
.map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?; .map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?;
WebDavAdapter::write_file_entry(&mut xml_writer, &file, &propfind_request, &base_href) WebDavAdapter::write_file_entry(
.map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?; &mut xml_writer,
&file,
&propfind_request,
&base_href,
)
.map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?;
WebDavAdapter::write_multistatus_end(&mut xml_writer) WebDavAdapter::write_multistatus_end(&mut xml_writer)
.map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?; .map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?;
} }
@@ -358,8 +361,13 @@ async fn handle_propfind(
let mut xml_writer = Writer::new(&mut buf); let mut xml_writer = Writer::new(&mut buf);
WebDavAdapter::write_multistatus_start(&mut xml_writer) WebDavAdapter::write_multistatus_start(&mut xml_writer)
.map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?; .map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?;
WebDavAdapter::write_file_entry(&mut xml_writer, &file, &propfind_request, &base_href) WebDavAdapter::write_file_entry(
.map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?; &mut xml_writer,
&file,
&propfind_request,
&base_href,
)
.map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?;
WebDavAdapter::write_multistatus_end(&mut xml_writer) WebDavAdapter::write_multistatus_end(&mut xml_writer)
.map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?; .map_err(|e| AppError::internal_error(format!("XML write error: {}", e)))?;
} }
@@ -910,13 +918,17 @@ async fn handle_delete(
folder_service folder_service
.delete_folder(&folder.id, caller_id) .delete_folder(&folder.id, caller_id)
.await .await
.map_err(|e| AppError::internal_error(format!("Failed to delete folder: {}", e)))?; .map_err(|e| {
AppError::internal_error(format!("Failed to delete folder: {}", e))
})?;
} }
Ok(ResolvedResource::File(file)) => { Ok(ResolvedResource::File(file)) => {
file_management_service file_management_service
.delete_file(&file.id) .delete_file(&file.id)
.await .await
.map_err(|e| AppError::internal_error(format!("Failed to delete file: {}", e)))?; .map_err(|e| {
AppError::internal_error(format!("Failed to delete file: {}", e))
})?;
} }
Err(_) => return Err(AppError::not_found(format!("Resource not found: {}", path))), Err(_) => return Err(AppError::not_found(format!("Resource not found: {}", path))),
} }
@@ -1002,8 +1014,14 @@ async fn handle_move(
let dest_exists = if let Some(resolver) = &state.path_resolver { let dest_exists = if let Some(resolver) = &state.path_resolver {
resolver.exists(&destination_path).await.unwrap_or(false) resolver.exists(&destination_path).await.unwrap_or(false)
} else { } else {
folder_service.get_folder_by_path(&destination_path).await.is_ok() folder_service
|| file_retrieval_service.get_file_by_path(&destination_path).await.is_ok() .get_folder_by_path(&destination_path)
.await
.is_ok()
|| file_retrieval_service
.get_file_by_path(&destination_path)
.await
.is_ok()
}; };
if dest_exists { if dest_exists {
return Err(AppError::precondition_failed( return Err(AppError::precondition_failed(
@@ -1044,7 +1062,9 @@ async fn handle_move(
folder.owner_id.as_deref().unwrap_or("webdav"), folder.owner_id.as_deref().unwrap_or("webdav"),
) )
.await .await
.map_err(|e| AppError::internal_error(format!("Failed to move folder: {}", e)))?; .map_err(|e| {
AppError::internal_error(format!("Failed to move folder: {}", e))
})?;
if folder.name != dest_folder_name { if folder.name != dest_folder_name {
let rename_dto = crate::application::dtos::folder_dto::RenameFolderDto { let rename_dto = crate::application::dtos::folder_dto::RenameFolderDto {
@@ -1057,7 +1077,9 @@ async fn handle_move(
folder.owner_id.as_deref().unwrap_or("webdav"), folder.owner_id.as_deref().unwrap_or("webdav"),
) )
.await .await
.map_err(|e| AppError::internal_error(format!("Failed to rename folder: {}", e)))?; .map_err(|e| {
AppError::internal_error(format!("Failed to rename folder: {}", e))
})?;
} }
} }
Ok(ResolvedResource::File(file)) => { Ok(ResolvedResource::File(file)) => {
@@ -1080,16 +1102,25 @@ async fn handle_move(
file_management_service file_management_service
.move_file(&file.id, Some(dest_parent_path.to_string())) .move_file(&file.id, Some(dest_parent_path.to_string()))
.await .await
.map_err(|e| AppError::internal_error(format!("Failed to move file: {}", e)))?; .map_err(|e| {
AppError::internal_error(format!("Failed to move file: {}", e))
})?;
} }
if file.name != dest_filename { if file.name != dest_filename {
file_management_service file_management_service
.rename_file(&file.id, dest_filename) .rename_file(&file.id, dest_filename)
.await .await
.map_err(|e| AppError::internal_error(format!("Failed to rename file: {}", e)))?; .map_err(|e| {
AppError::internal_error(format!("Failed to rename file: {}", e))
})?;
} }
} }
Err(_) => return Err(AppError::not_found(format!("Resource not found: {}", source_path))), Err(_) => {
return Err(AppError::not_found(format!(
"Resource not found: {}",
source_path
)));
}
} }
} else { } else {
// Fallback: legacy double-query path // Fallback: legacy double-query path
@@ -1137,13 +1168,17 @@ async fn handle_move(
folder.owner_id.as_deref().unwrap_or("webdav"), folder.owner_id.as_deref().unwrap_or("webdav"),
) )
.await .await
.map_err(|e| AppError::internal_error(format!("Failed to rename folder: {}", e)))?; .map_err(|e| {
AppError::internal_error(format!("Failed to rename folder: {}", e))
})?;
} }
} else { } else {
let file = file_retrieval_service let file = file_retrieval_service
.get_file_by_path(&source_path) .get_file_by_path(&source_path)
.await .await
.map_err(|_e| AppError::not_found(format!("Resource not found: {}", source_path)))?; .map_err(|_e| {
AppError::not_found(format!("Resource not found: {}", source_path))
})?;
let dest_filename = destination_path let dest_filename = destination_path
.split('/') .split('/')
@@ -1170,7 +1205,9 @@ async fn handle_move(
file_management_service file_management_service
.rename_file(&file.id, dest_filename) .rename_file(&file.id, dest_filename)
.await .await
.map_err(|e| AppError::internal_error(format!("Failed to rename file: {}", e)))?; .map_err(|e| {
AppError::internal_error(format!("Failed to rename file: {}", e))
})?;
} }
} }
} }
@@ -1240,8 +1277,14 @@ async fn handle_copy(
let dest_exists = if let Some(resolver) = &state.path_resolver { let dest_exists = if let Some(resolver) = &state.path_resolver {
resolver.exists(&destination_path).await.unwrap_or(false) resolver.exists(&destination_path).await.unwrap_or(false)
} else { } else {
folder_service.get_folder_by_path(&destination_path).await.is_ok() folder_service
|| file_retrieval_service.get_file_by_path(&destination_path).await.is_ok() .get_folder_by_path(&destination_path)
.await
.is_ok()
|| file_retrieval_service
.get_file_by_path(&destination_path)
.await
.is_ok()
}; };
if dest_exists { if dest_exists {
return Err(AppError::precondition_failed( return Err(AppError::precondition_failed(
@@ -1296,7 +1339,10 @@ async fn handle_copy(
.create_folder(create_dto) .create_folder(create_dto)
.await .await
.map_err(|e| { .map_err(|e| {
AppError::internal_error(format!("Failed to create destination folder: {}", e)) AppError::internal_error(format!(
"Failed to create destination folder: {}",
e
))
})?; })?;
} }
} }
@@ -1322,7 +1368,12 @@ async fn handle_copy(
.await .await
.map_err(|e| AppError::internal_error(format!("Failed to copy file: {}", e)))?; .map_err(|e| AppError::internal_error(format!("Failed to copy file: {}", e)))?;
} }
Err(_) => return Err(AppError::not_found(format!("Resource not found: {}", source_path))), Err(_) => {
return Err(AppError::not_found(format!(
"Resource not found: {}",
source_path
)));
}
} }
} else { } else {
// Fallback: legacy double-query path // Fallback: legacy double-query path
@@ -1371,14 +1422,19 @@ async fn handle_copy(
.create_folder(create_dto) .create_folder(create_dto)
.await .await
.map_err(|e| { .map_err(|e| {
AppError::internal_error(format!("Failed to create destination folder: {}", e)) AppError::internal_error(format!(
"Failed to create destination folder: {}",
e
))
})?; })?;
} }
} else { } else {
let file = file_retrieval_service let file = file_retrieval_service
.get_file_by_path(&source_path) .get_file_by_path(&source_path)
.await .await
.map_err(|_e| AppError::not_found(format!("Resource not found: {}", source_path)))?; .map_err(|_e| {
AppError::not_found(format!("Resource not found: {}", source_path))
})?;
let dest_parent_path = if let Some(idx) = destination_path.rfind('/') { let dest_parent_path = if let Some(idx) = destination_path.rfind('/') {
&destination_path[..idx] &destination_path[..idx]
+5 -9
View File
@@ -205,10 +205,7 @@ pub async fn auth_middleware(
} }
Err(e) => { Err(e) => {
tracing::warn!("Bearer token validation failed: {}", e); tracing::warn!("Bearer token validation failed: {}", e);
return Err(AuthError::InvalidToken(format!( return Err(AuthError::InvalidToken(format!("Invalid token: {}", e)));
"Invalid token: {}",
e
)));
} }
} }
} }
@@ -273,7 +270,9 @@ pub async fn auth_middleware(
{ {
use crate::interfaces::api::cookie_auth; use crate::interfaces::api::cookie_auth;
if let Some(token_str) = cookie_auth::extract_cookie_value(&headers, cookie_auth::ACCESS_COOKIE) { if let Some(token_str) =
cookie_auth::extract_cookie_value(&headers, cookie_auth::ACCESS_COOKIE)
{
if !token_str.is_empty() { if !token_str.is_empty() {
tracing::debug!("Processing cookie-based authentication"); tracing::debug!("Processing cookie-based authentication");
@@ -281,10 +280,7 @@ pub async fn auth_middleware(
let token_service = &auth_service.token_service; let token_service = &auth_service.token_service;
match token_service.validate_token(&token_str) { match token_service.validate_token(&token_str) {
Ok(claims) => { Ok(claims) => {
tracing::debug!( tracing::debug!("Cookie token validated for user: {}", claims.username);
"Cookie token validated for user: {}",
claims.username
);
let current_user = CurrentUser { let current_user = CurrentUser {
id: claims.sub, id: claims.sub,
username: claims.username, username: claims.username,
+73 -75
View File
@@ -1,75 +1,73 @@
//! CSRF double-submit cookie middleware. //! CSRF double-submit cookie middleware.
//! //!
//! State-changing requests (`POST`, `PUT`, `DELETE`, `PATCH`) that were //! State-changing requests (`POST`, `PUT`, `DELETE`, `PATCH`) that were
//! authenticated via an HttpOnly cookie (i.e. browser sessions) **must** //! authenticated via an HttpOnly cookie (i.e. browser sessions) **must**
//! include an `X-CSRF-Token` header whose value matches the `oxicloud_csrf` //! include an `X-CSRF-Token` header whose value matches the `oxicloud_csrf`
//! cookie. Requests authenticated via `Bearer` or `Basic` headers are //! cookie. Requests authenticated via `Bearer` or `Basic` headers are
//! exempt because they are not vulnerable to CSRF — the browser never //! exempt because they are not vulnerable to CSRF — the browser never
//! attaches those automatically. //! attaches those automatically.
//! //!
//! Safe methods (`GET`, `HEAD`, `OPTIONS`) are always allowed through. //! Safe methods (`GET`, `HEAD`, `OPTIONS`) are always allowed through.
use axum::{ use axum::{
extract::Request, extract::Request,
http::{Method, StatusCode}, http::{Method, StatusCode},
middleware::Next, middleware::Next,
response::{IntoResponse, Response}, response::{IntoResponse, Response},
}; };
use crate::interfaces::api::cookie_auth; use crate::interfaces::api::cookie_auth;
use crate::interfaces::middleware::auth::CookieAuthenticated; use crate::interfaces::middleware::auth::CookieAuthenticated;
/// Methods considered safe (no side-effects) — CSRF check is skipped. /// Methods considered safe (no side-effects) — CSRF check is skipped.
const SAFE_METHODS: [Method; 3] = [Method::GET, Method::HEAD, Method::OPTIONS]; const SAFE_METHODS: [Method; 3] = [Method::GET, Method::HEAD, Method::OPTIONS];
/// Middleware that enforces CSRF protection for cookie-authenticated browser /// Middleware that enforces CSRF protection for cookie-authenticated browser
/// sessions using the **double-submit cookie** pattern. /// sessions using the **double-submit cookie** pattern.
/// ///
/// Must be applied **after** `auth_middleware` so that the /// Must be applied **after** `auth_middleware` so that the
/// `CookieAuthenticated` marker is available in extensions. /// `CookieAuthenticated` marker is available in extensions.
pub async fn csrf_middleware(request: Request, next: Next) -> Result<Response, Response> { pub async fn csrf_middleware(request: Request, next: Next) -> Result<Response, Response> {
// Safe methods never need CSRF validation. // Safe methods never need CSRF validation.
if SAFE_METHODS.contains(request.method()) { if SAFE_METHODS.contains(request.method()) {
return Ok(next.run(request).await); return Ok(next.run(request).await);
} }
// Only enforce for cookie-authenticated sessions. // Only enforce for cookie-authenticated sessions.
let is_cookie_auth = request.extensions().get::<CookieAuthenticated>().is_some(); let is_cookie_auth = request.extensions().get::<CookieAuthenticated>().is_some();
if !is_cookie_auth { if !is_cookie_auth {
return Ok(next.run(request).await); return Ok(next.run(request).await);
} }
// Extract the CSRF token from the cookie. // Extract the CSRF token from the cookie.
let cookie_token = cookie_auth::extract_cookie_value( let cookie_token =
request.headers(), cookie_auth::extract_cookie_value(request.headers(), cookie_auth::CSRF_COOKIE);
cookie_auth::CSRF_COOKIE,
); // Extract the CSRF token from the request header.
let header_token = request
// Extract the CSRF token from the request header. .headers()
let header_token = request .get(cookie_auth::CSRF_HEADER)
.headers() .and_then(|v| v.to_str().ok())
.get(cookie_auth::CSRF_HEADER) .map(|s| s.to_string());
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string()); match (cookie_token, header_token) {
(Some(c), Some(h)) if !c.is_empty() && c == h => {
match (cookie_token, header_token) { // Tokens match — allow the request through.
(Some(c), Some(h)) if !c.is_empty() && c == h => { Ok(next.run(request).await)
// Tokens match — allow the request through. }
Ok(next.run(request).await) _ => {
} tracing::warn!(
_ => { method = %request.method(),
tracing::warn!( uri = %request.uri(),
method = %request.method(), "CSRF validation failed: missing or mismatched token"
uri = %request.uri(), );
"CSRF validation failed: missing or mismatched token" Err((
); StatusCode::FORBIDDEN,
Err(( axum::Json(serde_json::json!({
StatusCode::FORBIDDEN, "error": "CSRF token missing or invalid"
axum::Json(serde_json::json!({ })),
"error": "CSRF token missing or invalid" )
})), .into_response())
) }
.into_response()) }
} }
}
}
+196 -202
View File
@@ -1,202 +1,196 @@
//! IP-based rate limiting middleware for authentication endpoints. //! IP-based rate limiting middleware for authentication endpoints.
//! //!
//! Uses `moka` TTL caches (already a project dependency) to track request //! Uses `moka` TTL caches (already a project dependency) to track request
//! counts per client IP. Each protected endpoint group gets its own //! counts per client IP. Each protected endpoint group gets its own
//! [`RateLimiter`] instance with independently tuneable limits. //! [`RateLimiter`] instance with independently tuneable limits.
//! //!
//! The middleware extracts the client IP from (in order): //! The middleware extracts the client IP from (in order):
//! 1. `X-Forwarded-For` header (first entry — set by reverse proxies) //! 1. `X-Forwarded-For` header (first entry — set by reverse proxies)
//! 2. `X-Real-Ip` header //! 2. `X-Real-Ip` header
//! 3. The TCP peer address from the connection info //! 3. The TCP peer address from the connection info
//! //!
//! When the limit is exceeded a `429 Too Many Requests` response is returned //! When the limit is exceeded a `429 Too Many Requests` response is returned
//! with a `Retry-After` header indicating how many seconds to wait. //! with a `Retry-After` header indicating how many seconds to wait.
use axum::{ use axum::{
extract::ConnectInfo, extract::ConnectInfo,
http::{HeaderValue, Request, StatusCode}, http::{HeaderValue, Request, StatusCode},
middleware::Next, middleware::Next,
response::{IntoResponse, Response}, response::{IntoResponse, Response},
}; };
use moka::sync::Cache; use moka::sync::Cache;
use std::net::SocketAddr; use std::net::SocketAddr;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
/// A simple sliding-window counter keyed by IP address. /// A simple sliding-window counter keyed by IP address.
/// ///
/// Each key lives for `window` seconds; every request increments the counter. /// Each key lives for `window` seconds; every request increments the counter.
/// Once the counter reaches `max_requests` the request is rejected. /// Once the counter reaches `max_requests` the request is rejected.
#[derive(Clone)] #[derive(Clone)]
pub struct RateLimiter { pub struct RateLimiter {
/// Maps `IP -> request_count` with automatic TTL expiration. /// Maps `IP -> request_count` with automatic TTL expiration.
cache: Cache<String, u32>, cache: Cache<String, u32>,
/// Maximum requests allowed within the window. /// Maximum requests allowed within the window.
max_requests: u32, max_requests: u32,
/// Window duration in seconds (also used for `Retry-After`). /// Window duration in seconds (also used for `Retry-After`).
window_secs: u64, window_secs: u64,
} }
impl RateLimiter { impl RateLimiter {
/// Create a new rate limiter. /// Create a new rate limiter.
/// ///
/// * `max_requests` — ceiling per IP within the window /// * `max_requests` — ceiling per IP within the window
/// * `window_secs` — sliding window duration /// * `window_secs` — sliding window duration
/// * `max_entries` — upper bound on tracked IPs (evicts LRU when exceeded) /// * `max_entries` — upper bound on tracked IPs (evicts LRU when exceeded)
pub fn new(max_requests: u32, window_secs: u64, max_entries: u64) -> Self { pub fn new(max_requests: u32, window_secs: u64, max_entries: u64) -> Self {
let cache = Cache::builder() let cache = Cache::builder()
.time_to_live(Duration::from_secs(window_secs)) .time_to_live(Duration::from_secs(window_secs))
.max_capacity(max_entries) .max_capacity(max_entries)
.build(); .build();
Self { Self {
cache, cache,
max_requests, max_requests,
window_secs, window_secs,
} }
} }
/// Check whether the IP is allowed. Returns `Ok(current_count)` or /// Check whether the IP is allowed. Returns `Ok(current_count)` or
/// `Err(StatusCode::TOO_MANY_REQUESTS)`. /// `Err(StatusCode::TOO_MANY_REQUESTS)`.
pub fn check_and_increment(&self, ip: &str) -> Result<u32, ()> { pub fn check_and_increment(&self, ip: &str) -> Result<u32, ()> {
let key = ip.to_string(); let key = ip.to_string();
// moka's entry API lets us atomically read-modify-write. // moka's entry API lets us atomically read-modify-write.
// On first access the entry is inserted with count = 1 and the TTL // On first access the entry is inserted with count = 1 and the TTL
// starts. Subsequent accesses within the window increment the count. // starts. Subsequent accesses within the window increment the count.
let count = self let count = self.cache.entry(key).or_insert_with(|| 0).into_value() + 1;
.cache
.entry(key) // Write back the incremented value. Because `or_insert_with` returns
.or_insert_with(|| 0) // the *existing* value when the key was already present, we must always
.into_value() // re-insert so the counter actually advances. The TTL of the **first**
+ 1; // insert still governs eviction because moka uses insert-time TTL.
// However, on re-insert moka resets the TTL — for rate limiting this
// Write back the incremented value. Because `or_insert_with` returns // is fine because it means the window "slides" forward on activity.
// the *existing* value when the key was already present, we must always self.cache.insert(ip.to_string(), count);
// re-insert so the counter actually advances. The TTL of the **first**
// insert still governs eviction because moka uses insert-time TTL. if count > self.max_requests {
// However, on re-insert moka resets the TTL — for rate limiting this Err(())
// is fine because it means the window "slides" forward on activity. } else {
self.cache Ok(count)
.insert(ip.to_string(), count); }
}
if count > self.max_requests {
Err(()) /// Seconds the client should wait before retrying.
} else { pub fn retry_after(&self) -> u64 {
Ok(count) self.window_secs
} }
} }
/// Seconds the client should wait before retrying. // ─── Axum middleware factories ──────────────────────────────────────────────
pub fn retry_after(&self) -> u64 {
self.window_secs /// Extract the most-likely real client IP from headers / connection info.
} pub fn extract_client_ip<B>(req: &Request<B>) -> String {
} let headers = req.headers();
// ─── Axum middleware factories ────────────────────────────────────────────── // 1. X-Forwarded-For (first entry — closest to the client)
if let Some(xff) = headers.get("x-forwarded-for").and_then(|v| v.to_str().ok()) {
/// Extract the most-likely real client IP from headers / connection info. if let Some(first) = xff.split(',').next() {
pub fn extract_client_ip<B>(req: &Request<B>) -> String { let ip = first.trim();
let headers = req.headers(); if !ip.is_empty() {
return ip.to_string();
// 1. X-Forwarded-For (first entry — closest to the client) }
if let Some(xff) = headers.get("x-forwarded-for").and_then(|v| v.to_str().ok()) { }
if let Some(first) = xff.split(',').next() { }
let ip = first.trim();
if !ip.is_empty() { // 2. X-Real-Ip
return ip.to_string(); if let Some(xri) = headers.get("x-real-ip").and_then(|v| v.to_str().ok()) {
} let ip = xri.trim();
} if !ip.is_empty() {
} return ip.to_string();
}
// 2. X-Real-Ip }
if let Some(xri) = headers.get("x-real-ip").and_then(|v| v.to_str().ok()) {
let ip = xri.trim(); // 3. TCP peer (ConnectInfo extension set by axum::serve)
if !ip.is_empty() { if let Some(addr) = req.extensions().get::<ConnectInfo<SocketAddr>>() {
return ip.to_string(); return addr.0.ip().to_string();
} }
}
// Fallback — should never happen behind axum::serve
// 3. TCP peer (ConnectInfo extension set by axum::serve) "unknown".to_string()
if let Some(addr) = req.extensions().get::<ConnectInfo<SocketAddr>>() { }
return addr.0.ip().to_string();
} /// Build a rate-limit response with the standard `Retry-After` header.
fn too_many_requests(retry_after: u64) -> Response {
// Fallback — should never happen behind axum::serve let body = serde_json::json!({
"unknown".to_string() "error": "Too many requests",
} "retry_after_secs": retry_after,
});
/// Build a rate-limit response with the standard `Retry-After` header. let mut resp = (StatusCode::TOO_MANY_REQUESTS, axum::Json(body)).into_response();
fn too_many_requests(retry_after: u64) -> Response { if let Ok(val) = HeaderValue::from_str(&retry_after.to_string()) {
let body = serde_json::json!({ resp.headers_mut().insert("retry-after", val);
"error": "Too many requests", }
"retry_after_secs": retry_after, resp
}); }
let mut resp = (StatusCode::TOO_MANY_REQUESTS, axum::Json(body)).into_response();
if let Ok(val) = HeaderValue::from_str(&retry_after.to_string()) { /// Axum middleware: rate-limit login attempts.
resp.headers_mut().insert("retry-after", val); ///
} /// Inject via:
resp /// ```ignore
} /// .layer(axum::middleware::from_fn_with_state(limiter, rate_limit_login))
/// ```
/// Axum middleware: rate-limit login attempts. pub async fn rate_limit_login(
/// State(limiter): axum::extract::State<Arc<RateLimiter>>,
/// Inject via: req: Request<axum::body::Body>,
/// ```ignore next: Next,
/// .layer(axum::middleware::from_fn_with_state(limiter, rate_limit_login)) ) -> Response {
/// ``` let ip = extract_client_ip(&req);
pub async fn rate_limit_login( match limiter.check_and_increment(&ip) {
State(limiter): axum::extract::State<Arc<RateLimiter>>, Ok(_) => next.run(req).await,
req: Request<axum::body::Body>, Err(()) => {
next: Next, tracing::warn!(
) -> Response { ip = %ip,
let ip = extract_client_ip(&req); "Rate limit exceeded on login endpoint"
match limiter.check_and_increment(&ip) { );
Ok(_) => next.run(req).await, too_many_requests(limiter.retry_after())
Err(()) => { }
tracing::warn!( }
ip = %ip, }
"Rate limit exceeded on login endpoint"
); /// Axum middleware: rate-limit registration attempts.
too_many_requests(limiter.retry_after()) pub async fn rate_limit_register(
} State(limiter): axum::extract::State<Arc<RateLimiter>>,
} req: Request<axum::body::Body>,
} next: Next,
) -> Response {
/// Axum middleware: rate-limit registration attempts. let ip = extract_client_ip(&req);
pub async fn rate_limit_register( match limiter.check_and_increment(&ip) {
State(limiter): axum::extract::State<Arc<RateLimiter>>, Ok(_) => next.run(req).await,
req: Request<axum::body::Body>, Err(()) => {
next: Next, tracing::warn!(
) -> Response { ip = %ip,
let ip = extract_client_ip(&req); "Rate limit exceeded on register endpoint"
match limiter.check_and_increment(&ip) { );
Ok(_) => next.run(req).await, too_many_requests(limiter.retry_after())
Err(()) => { }
tracing::warn!( }
ip = %ip, }
"Rate limit exceeded on register endpoint"
); /// Axum middleware: rate-limit token refresh attempts.
too_many_requests(limiter.retry_after()) pub async fn rate_limit_refresh(
} State(limiter): axum::extract::State<Arc<RateLimiter>>,
} req: Request<axum::body::Body>,
} next: Next,
) -> Response {
/// Axum middleware: rate-limit token refresh attempts. let ip = extract_client_ip(&req);
pub async fn rate_limit_refresh( match limiter.check_and_increment(&ip) {
State(limiter): axum::extract::State<Arc<RateLimiter>>, Ok(_) => next.run(req).await,
req: Request<axum::body::Body>, Err(()) => {
next: Next, tracing::warn!(
) -> Response { ip = %ip,
let ip = extract_client_ip(&req); "Rate limit exceeded on refresh endpoint"
match limiter.check_and_increment(&ip) { );
Ok(_) => next.run(req).await, too_many_requests(limiter.retry_after())
Err(()) => { }
tracing::warn!( }
ip = %ip, }
"Rate limit exceeded on refresh endpoint"
); use axum::extract::State;
too_many_requests(limiter.retry_after())
}
}
}
use axum::extract::State;
+27 -13
View File
@@ -170,13 +170,15 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
); );
} }
if config.features.enable_auth { if config.features.enable_auth {
use interfaces::api::handlers::auth_handler::{auth_routes, login_route, register_route, refresh_route}; use interfaces::api::handlers::auth_handler::{
use oxicloud::interfaces::api::handlers::device_auth_handler; auth_routes, login_route, refresh_route, register_route,
};
use oxicloud::interfaces::api::handlers::app_password_handler; use oxicloud::interfaces::api::handlers::app_password_handler;
use oxicloud::interfaces::api::handlers::device_auth_handler;
use oxicloud::interfaces::middleware::auth::auth_middleware; use oxicloud::interfaces::middleware::auth::auth_middleware;
use oxicloud::interfaces::middleware::csrf::csrf_middleware; use oxicloud::interfaces::middleware::csrf::csrf_middleware;
use oxicloud::interfaces::middleware::rate_limit::{ use oxicloud::interfaces::middleware::rate_limit::{
RateLimiter, rate_limit_login, rate_limit_register, rate_limit_refresh, RateLimiter, rate_limit_login, rate_limit_refresh, rate_limit_register,
}; };
// ── Rate limiters (IP-based, in-memory via moka) ──────────────── // ── Rate limiters (IP-based, in-memory via moka) ────────────────
@@ -198,28 +200,40 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
)); ));
tracing::info!( tracing::info!(
"Rate limiting enabled — login: {}/{} s, register: {}/{} s, refresh: {}/{} s", "Rate limiting enabled — login: {}/{} s, register: {}/{} s, refresh: {}/{} s",
rl.login_max_requests, rl.login_window_secs, rl.login_max_requests,
rl.register_max_requests, rl.register_window_secs, rl.login_window_secs,
rl.refresh_max_requests, rl.refresh_window_secs, rl.register_max_requests,
rl.register_window_secs,
rl.refresh_max_requests,
rl.refresh_window_secs,
); );
// Auth routes split by rate-limit policy // Auth routes split by rate-limit policy
let auth_login = login_route() let auth_login = login_route()
.layer(axum::middleware::from_fn_with_state(login_limiter.clone(), rate_limit_login)) .layer(axum::middleware::from_fn_with_state(
login_limiter.clone(),
rate_limit_login,
))
.with_state(app_state.clone()); .with_state(app_state.clone());
let auth_register = register_route() let auth_register = register_route()
.layer(axum::middleware::from_fn_with_state(register_limiter.clone(), rate_limit_register)) .layer(axum::middleware::from_fn_with_state(
register_limiter.clone(),
rate_limit_register,
))
.with_state(app_state.clone()); .with_state(app_state.clone());
let auth_refresh = refresh_route() let auth_refresh = refresh_route()
.layer(axum::middleware::from_fn_with_state(refresh_limiter.clone(), rate_limit_refresh)) .layer(axum::middleware::from_fn_with_state(
refresh_limiter.clone(),
rate_limit_refresh,
))
.with_state(app_state.clone()); .with_state(app_state.clone());
// Remaining auth routes (status, OIDC, protected /me, /logout, etc.) // Remaining auth routes (status, OIDC, protected /me, /logout, etc.)
let auth_router = auth_routes().with_state(app_state.clone()); let auth_router = auth_routes().with_state(app_state.clone());
// Device Authorization Grant (RFC 8628) // Device Authorization Grant (RFC 8628)
// Public endpoints: /api/auth/device/authorize + /api/auth/device/token // Public endpoints: /api/auth/device/authorize + /api/auth/device/token
let device_public = device_auth_handler::device_auth_public_routes() let device_public =
.with_state(app_state.clone()); device_auth_handler::device_auth_public_routes().with_state(app_state.clone());
// Protected endpoints: /api/auth/device/verify, /api/auth/device/devices // Protected endpoints: /api/auth/device/verify, /api/auth/device/devices
let device_protected = device_auth_handler::device_auth_protected_routes() let device_protected = device_auth_handler::device_auth_protected_routes()
.layer(axum::middleware::from_fn(csrf_middleware)) .layer(axum::middleware::from_fn(csrf_middleware))
@@ -325,8 +339,8 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
// ── Security headers ───────────────────────────────────────────────── // ── Security headers ─────────────────────────────────────────────────
// Applied globally so every response (API, static, DAV) carries them. // Applied globally so every response (API, static, DAV) carries them.
use axum::http::header::HeaderName;
use axum::http::HeaderValue; use axum::http::HeaderValue;
use axum::http::header::HeaderName;
app = app app = app
.layer(SetResponseHeaderLayer::overriding( .layer(SetResponseHeaderLayer::overriding(
@@ -347,7 +361,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
frame-src *; \ frame-src *; \
frame-ancestors 'none'; \ frame-ancestors 'none'; \
base-uri 'self'; \ base-uri 'self'; \
form-action 'self'" form-action 'self'",
), ),
)) ))
.layer(SetResponseHeaderLayer::overriding( .layer(SetResponseHeaderLayer::overriding(