use std::net::IpAddr;
use std::sync::Arc;
use axess_clock::Clock;
use axum::http::request::Parts;
use crate::authn::ids::{DeviceId, TenantId, UserId};
use crate::device::fingerprint::DeviceFingerprintExtractor;
use crate::device::lifecycle::DeviceLifecycleService;
use crate::device::resolver::DeviceResolver;
use crate::device::store::DeviceStore;
type TenantFn = Arc<dyn Fn(&Parts) -> Option<TenantId> + Send + Sync>;
type ClientIpFn = Arc<dyn Fn(&Parts) -> Option<IpAddr> + Send + Sync>;
type UserFn = Arc<dyn Fn(&Parts) -> Option<UserId> + Send + Sync>;
type NewIdFn = Arc<dyn Fn() -> DeviceId + Send + Sync>;
pub struct LifecycleDeviceResolver<E, S, C>
where
E: DeviceFingerprintExtractor,
S: DeviceStore,
C: Clock,
{
extractor: Arc<E>,
lifecycle: DeviceLifecycleService<S>,
clock: Arc<C>,
tenant_fn: TenantFn,
client_ip_fn: ClientIpFn,
user_fn: UserFn,
new_id_fn: NewIdFn,
}
impl<E, S, C> Clone for LifecycleDeviceResolver<E, S, C>
where
E: DeviceFingerprintExtractor,
S: DeviceStore,
C: Clock,
{
fn clone(&self) -> Self {
Self {
extractor: self.extractor.clone(),
lifecycle: self.lifecycle.clone(),
clock: self.clock.clone(),
tenant_fn: self.tenant_fn.clone(),
client_ip_fn: self.client_ip_fn.clone(),
user_fn: self.user_fn.clone(),
new_id_fn: self.new_id_fn.clone(),
}
}
}
impl<E, S, C> LifecycleDeviceResolver<E, S, C>
where
E: DeviceFingerprintExtractor,
S: DeviceStore,
C: Clock,
{
pub fn new(extractor: E, lifecycle: DeviceLifecycleService<S>, clock: C) -> Self {
Self {
extractor: Arc::new(extractor),
lifecycle,
clock: Arc::new(clock),
tenant_fn: default_tenant_fn(),
client_ip_fn: default_client_ip_fn(),
user_fn: default_user_fn(),
new_id_fn: default_new_id_fn(),
}
}
pub fn with_tenant_fn<F>(mut self, f: F) -> Self
where
F: Fn(&Parts) -> Option<TenantId> + Send + Sync + 'static,
{
self.tenant_fn = Arc::new(f);
self
}
pub fn with_client_ip_fn<F>(mut self, f: F) -> Self
where
F: Fn(&Parts) -> Option<IpAddr> + Send + Sync + 'static,
{
self.client_ip_fn = Arc::new(f);
self
}
pub fn with_user_fn<F>(mut self, f: F) -> Self
where
F: Fn(&Parts) -> Option<UserId> + Send + Sync + 'static,
{
self.user_fn = Arc::new(f);
self
}
pub fn with_new_id_fn<F>(mut self, f: F) -> Self
where
F: Fn() -> DeviceId + Send + Sync + 'static,
{
self.new_id_fn = Arc::new(f);
self
}
}
impl<E, S, C> DeviceResolver for LifecycleDeviceResolver<E, S, C>
where
E: DeviceFingerprintExtractor,
S: DeviceStore,
C: Clock,
{
type Error = S::Error;
async fn resolve(&self, parts: &Parts) -> Result<Option<DeviceId>, Self::Error> {
let Some(tenant) = (self.tenant_fn)(parts) else {
return Ok(None);
};
let client_ip = (self.client_ip_fn)(parts);
let Some(fp) = self.extractor.extract(&tenant, parts, client_ip) else {
return Ok(None);
};
let user = (self.user_fn)(parts);
let now = self.clock.now();
let new_id_fn = self.new_id_fn.clone();
let id = self
.lifecycle
.ensure_device(&tenant, user.as_ref(), fp, now, move || (new_id_fn)())
.await?;
Ok(Some(id))
}
}
fn default_tenant_fn() -> TenantFn {
Arc::new(|parts: &Parts| parts.extensions.get::<TenantId>().cloned())
}
fn default_client_ip_fn() -> ClientIpFn {
Arc::new(|parts: &Parts| {
parts
.extensions
.get::<crate::client_ip::ClientIp>()
.copied()
.unwrap_or_default()
.get()
})
}
fn default_user_fn() -> UserFn {
Arc::new(|parts: &Parts| {
parts.uri.host();
None
})
}
fn default_new_id_fn() -> NewIdFn {
Arc::new(|| {
DeviceId::try_new(uuid::Uuid::new_v4().to_string())
.expect("uuid::Uuid::new_v4() always produces a DeviceId-valid string")
})
}
#[cfg(test)]
mod tests;