use std::future::Future;
use std::sync::LazyLock;
use argon2::Argon2;
use argon2::password_hash::phc::PasswordHash;
use argon2::password_hash::{PasswordHasher, PasswordVerifier};
use serde::{Deserialize, Serialize};
use super::Policy;
use crate::db::{DateTime, Db, DbValue, Executor, FromRow, Model, Row, ToDbValue, sql};
use crate::{Error, Result};
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[non_exhaustive]
pub struct User {
pub id: i64,
pub name: String,
pub email: String,
#[serde(skip_serializing, default)]
pub password: String,
pub email_verified_at: Option<DateTime>,
pub created_at: Option<DateTime>,
pub updated_at: Option<DateTime>,
#[serde(flatten, default)]
pub extra: std::collections::BTreeMap<String, serde_json::Value>,
}
const NOT_EXTRA: &[&str] = &[
"id",
"name",
"email",
"password",
"email_verified_at",
"created_at",
"updated_at",
"sessions_revoked_at",
"remember_token",
"session_revoked",
];
impl FromRow for User {
fn from_row(row: &Row) -> std::result::Result<Self, crate::db::DbError> {
let extra = row
.columns()
.into_iter()
.filter(|column| !NOT_EXTRA.contains(column))
.map(|column| (column.to_owned(), row.json(column)))
.collect();
Ok(Self {
id: row.try_get("id")?,
name: row.try_get("name")?,
email: row.try_get("email")?,
password: row.try_get("password")?,
email_verified_at: row.try_get("email_verified_at")?,
created_at: row.try_get("created_at")?,
updated_at: row.try_get("updated_at")?,
extra,
})
}
}
impl Model for User {
const TABLE: &'static str = "users";
const SELECT_ALL: bool = true;
const COLUMNS: &'static [&'static str] = &[
"id",
"name",
"email",
"password",
"email_verified_at",
"created_at",
"updated_at",
];
type Key = i64;
fn id(&self) -> i64 {
self.id
}
fn set_id(&mut self, id: i64) {
self.id = id;
}
fn values(&self) -> Vec<DbValue> {
vec![
self.name.to_db_value(),
self.email.to_db_value(),
self.password.to_db_value(),
self.email_verified_at.to_db_value(),
self.created_at.to_db_value(),
self.updated_at.to_db_value(),
]
}
fn touch(&mut self, now: DateTime, creating: bool) {
if creating && self.created_at.is_none() {
self.created_at = Some(now);
}
self.updated_at = Some(now);
}
}
impl User {
pub fn find_by_email<'c, E: Executor<'c>>(
db: E,
email: &str,
) -> impl Future<Output = Result<Option<Self>>> + Send {
Self::query()
.where_eq("email", normalize_email(email))
.first(db)
}
pub async fn register(db: &Db, name: &str, email: &str, password: &str) -> Result<Self> {
let user = Self {
name: name.trim().to_owned(),
email: normalize_email(email),
password: hash_password(password).await?,
..Self::default()
};
let id = Self::create(db, user).await?.id;
Self::find_or_404(db, id).await
}
pub async fn set_password(&mut self, db: &Db, password: &str) -> Result {
self.password = hash_password(password).await?;
self.save(db).await
}
pub fn get<T: serde::de::DeserializeOwned>(&self, column: &str) -> Option<T> {
let value = self.extra.get(column)?;
serde_json::from_value(value.clone()).ok().or_else(|| {
let flag = value.as_i64().filter(|n| *n == 0 || *n == 1)?;
serde_json::from_value(serde_json::Value::Bool(flag == 1)).ok()
})
}
pub async fn set(&mut self, db: &Db, column: &str, value: impl ToDbValue) -> Result {
let plain = !column.is_empty()
&& column
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_');
if !plain || NOT_EXTRA.contains(&column) {
return Err(anyhow::anyhow!("User::set can't change `{column}`").into());
}
let value = value.to_db_value();
sql(format!(
"UPDATE users SET {} = ? WHERE id = ?",
crate::db::quote(column)
))
.bind(value.clone())
.bind(self.id)
.execute(db)
.await?;
self.extra.insert(column.to_owned(), value.to_json());
Ok(())
}
pub async fn revoke_sessions(&self, db: &Db) -> Result {
revoke_sessions(db, self.id).await.map(|_| ())
}
pub(crate) async fn find_with_revocation(
db: &Db,
id: i64,
session_id: &str,
) -> Result<Option<(Self, i64, bool)>> {
let row = sql(
"SELECT users.*, (SELECT COUNT(*) FROM revoked_sessions WHERE id = ?) AS session_revoked \
FROM users WHERE id = ?",
)
.bind(session_id)
.bind(id)
.fetch_optional(db)
.await?;
Ok(match row {
Some(row) => {
let revoked: i64 = row.try_get("session_revoked")?;
Some((
Self::from_row(&row)?,
row.try_get("sessions_revoked_at")?,
revoked > 0,
))
}
None => None,
})
}
pub async fn attempt(db: &Db, email: &str, password: &str) -> Result<Option<Self>> {
let user = Self::find_by_email(db, email).await?;
let hash = user
.as_ref()
.map_or_else(dummy_hash, |u| u.password.clone());
let valid = verify_password(password, &hash).await;
let Some(mut user) = user.filter(|_| valid) else {
return Ok(None);
};
user.rehash_if_needed(db, password).await?;
Ok(Some(user))
}
pub(crate) async fn rehash_if_needed(&mut self, db: &Db, password: &str) -> Result {
if needs_rehash(&self.password) {
self.set_password(db, password).await?;
}
Ok(())
}
pub fn has_password(&self) -> bool {
!self.password.is_empty()
}
pub async fn check_password(&self, password: &str) -> bool {
verify_password(password, &self.password).await
}
pub fn can(&self, ability: &str, target: &impl Policy) -> bool {
target.allows(self, ability)
}
pub fn authorize(&self, ability: &str, target: &impl Policy) -> Result {
if self.can(ability, target) {
Ok(())
} else {
Err(Error::Forbidden)
}
}
}
pub(crate) async fn revoke_sessions(db: &Db, id: i64) -> Result<i64> {
let now = super::unix_millis();
sql("UPDATE users SET sessions_revoked_at = ? WHERE id = ?")
.bind(now)
.bind(id)
.execute(db)
.await?;
Ok(now)
}
pub(crate) fn normalize_email(email: &str) -> String {
email.trim().to_lowercase()
}
pub(crate) fn dummy_hash() -> String {
static HASH: LazyLock<String> = LazyLock::new(|| {
Argon2::default()
.hash_password(b"renox-timing-equaliser")
.map(|hash| hash.to_string())
.unwrap_or_default()
});
HASH.clone()
}
pub async fn hash_password(password: &str) -> Result<String> {
let password = password.to_owned();
tokio::task::spawn_blocking(move || {
Argon2::default()
.hash_password(password.as_bytes())
.map(|hash| hash.to_string())
.map_err(|err| anyhow::anyhow!("could not hash the password: {err}"))
})
.await
.map_err(anyhow::Error::from)?
.map_err(Error::from)
}
pub async fn verify_password(password: &str, hash: &str) -> bool {
let (password, hash) = (password.to_owned(), hash.to_owned());
tokio::task::spawn_blocking(move || {
if is_bcrypt(&hash) {
let hash = hash.replacen("$2y$", "$2b$", 1);
return bcrypt::verify(password.as_bytes(), &hash).unwrap_or(false);
}
PasswordHash::new(&hash).is_ok_and(|parsed| {
Argon2::default()
.verify_password(password.as_bytes(), &parsed)
.is_ok()
})
})
.await
.unwrap_or(false)
}
fn is_bcrypt(hash: &str) -> bool {
["$2y$", "$2b$", "$2a$"].iter().any(|p| hash.starts_with(p))
}
pub fn needs_rehash(hash: &str) -> bool {
!hash.starts_with("$argon2id$")
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn hashes_verify() {
let hash = hash_password("rahasia123").await.unwrap();
assert!(hash.starts_with("$argon2id$"));
assert!(verify_password("rahasia123", &hash).await);
assert!(!verify_password("salah", &hash).await);
assert!(!verify_password("rahasia123", "not a hash").await);
assert!(!needs_rehash(&hash));
}
#[tokio::test]
async fn laravel_bcrypt_hashes_verify() {
let laravel = bcrypt::hash("password", 4)
.unwrap()
.replacen("$2b$", "$2y$", 1);
assert!(laravel.starts_with("$2y$04$"));
assert!(verify_password("password", &laravel).await);
assert!(!verify_password("wrong", &laravel).await);
assert!(needs_rehash(&laravel));
}
}