use std::collections::HashSet;
use async_trait::async_trait;
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use uuid::Uuid;
use crate::{
audit::logger::{AuditEventType, SecretType, get_audit_logger},
error::{AuthError, Result},
};
mod postgres;
pub use postgres::{PostgresAccountStore, SCHEMA_SQL};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ProviderLink {
pub provider: String,
pub provider_id: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AccountRecord {
pub user_id: String,
pub email: Option<String>,
pub providers: Vec<ProviderLink>,
}
#[async_trait]
pub trait AccountStore: Send + Sync {
async fn link_or_create_user(
&self,
email: Option<&str>,
email_verified: bool,
provider: &str,
provider_id: &str,
) -> Result<AccountLinkResult>;
async fn get_account(&self, user_id: &str) -> Result<AccountRecord>;
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AccountLinkResult {
pub user_id: String,
pub is_new: bool,
pub linked: bool,
}
pub struct InMemoryAccountStore {
by_identity: DashMap<String, String>,
by_user_id: DashMap<String, AccountRecord>,
}
impl InMemoryAccountStore {
#[must_use]
pub fn new() -> Self {
Self {
by_identity: DashMap::new(),
by_user_id: DashMap::new(),
}
}
#[must_use]
pub fn len(&self) -> usize {
self.by_user_id.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.by_user_id.is_empty()
}
}
impl Default for InMemoryAccountStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl AccountStore for InMemoryAccountStore {
async fn link_or_create_user(
&self,
email: Option<&str>,
email_verified: bool,
provider: &str,
provider_id: &str,
) -> Result<AccountLinkResult> {
let logger = get_audit_logger();
let verified_email = email.map(normalize_email).filter(|e| !e.is_empty() && email_verified);
let key = identity_key(verified_email.as_deref(), provider, provider_id);
let new_link = ProviderLink {
provider: provider.to_string(),
provider_id: provider_id.to_string(),
};
if let Some(existing_user_id) = self.by_identity.get(&key).map(|r| r.clone()) {
let mut record = self.by_user_id.get_mut(&existing_user_id).ok_or_else(|| {
AuthError::DatabaseError {
message: format!(
"account store inconsistency: identity '{key}' maps to missing user_id \
'{existing_user_id}'"
),
}
})?;
let already_linked = record.providers.contains(&new_link);
if !already_linked {
record.providers.push(new_link);
logger.log_success(
AuditEventType::AuthSuccess,
SecretType::SessionToken,
Some(existing_user_id.clone()),
&format!("account_linked:{provider}"),
);
}
return Ok(AccountLinkResult {
user_id: existing_user_id.clone(),
is_new: false,
linked: !already_linked,
});
}
let user_id = format!("user_{}", Uuid::new_v4().as_simple());
let record = AccountRecord {
user_id: user_id.clone(),
email: verified_email,
providers: vec![new_link],
};
self.by_identity.insert(key, user_id.clone());
self.by_user_id.insert(user_id.clone(), record);
logger.log_success(
AuditEventType::SessionTokenCreated,
SecretType::SessionToken,
Some(user_id.clone()),
&format!("account_created:{provider}"),
);
Ok(AccountLinkResult {
user_id,
is_new: true,
linked: false,
})
}
async fn get_account(&self, user_id: &str) -> Result<AccountRecord> {
self.by_user_id.get(user_id).map(|r| r.clone()).ok_or(AuthError::TokenNotFound)
}
}
#[must_use]
pub fn normalize_email(email: &str) -> String {
email.trim().to_lowercase()
}
fn identity_key(verified_email: Option<&str>, provider: &str, provider_id: &str) -> String {
match verified_email {
Some(email) => format!("email:{email}"),
None => format!("provider:{provider}\u{1f}{provider_id}"),
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TrustedEmailProviders {
providers: HashSet<String>,
}
impl TrustedEmailProviders {
#[must_use]
pub fn builtin_default() -> Self {
Self::only(["google", "apple"])
}
#[must_use]
pub fn none() -> Self {
Self {
providers: HashSet::new(),
}
}
#[must_use]
pub fn only(providers: impl IntoIterator<Item = impl Into<String>>) -> Self {
Self {
providers: providers.into_iter().map(|p| normalize_provider(&p.into())).collect(),
}
}
#[must_use]
pub fn trust(mut self, provider: impl Into<String>) -> Self {
self.providers.insert(normalize_provider(&provider.into()));
self
}
#[must_use]
pub fn distrust(mut self, provider: &str) -> Self {
self.providers.remove(&normalize_provider(provider));
self
}
#[must_use]
pub fn is_trusted(&self, provider: &str) -> bool {
self.providers.contains(&normalize_provider(provider))
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.providers.is_empty()
}
}
impl Default for TrustedEmailProviders {
fn default() -> Self {
Self::builtin_default()
}
}
fn normalize_provider(provider: &str) -> String {
provider.trim().to_lowercase()
}
#[allow(clippy::unwrap_used)] #[cfg(test)]
mod tests;