style: cargo fmt --all
This commit is contained in:
@@ -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/" {
|
||||||
|
|||||||
@@ -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!(
|
||||||
|
|||||||
@@ -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",
|
|
||||||
))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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,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),
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -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,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
@@ -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);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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"))
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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())
|
}
|
||||||
}
|
}
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user