mirror of
https://github.com/RaisFast/raisfast.git
synced 2026-09-30 00:03:21 +00:00
322 lines
8.9 KiB
Rust
322 lines
8.9 KiB
Rust
//! 短信验证码模型与数据库查询
|
|
|
|
use sqlx::FromRow;
|
|
|
|
use crate::db::dialect::ph;
|
|
use crate::errors::app_error::AppResult;
|
|
use crate::utils::id;
|
|
use crate::utils::tz::Timestamp;
|
|
|
|
/// 短信验证码数据库行模型
|
|
#[derive(Debug, FromRow)]
|
|
#[non_exhaustive]
|
|
pub struct SmsCode {
|
|
pub id: i64,
|
|
pub document_id: String,
|
|
pub phone: String,
|
|
pub code: String,
|
|
pub purpose: String,
|
|
pub expires_at: Timestamp,
|
|
pub verified_at: Option<Timestamp>,
|
|
pub attempts: i64,
|
|
pub ip_address: Option<String>,
|
|
pub created_at: Timestamp,
|
|
}
|
|
|
|
/// 生成指定位数的随机数字验证码
|
|
pub fn generate_code(length: u32) -> String {
|
|
let digits: Vec<u8> = (0..length)
|
|
.map(|_| {
|
|
let mut byte = [0u8; 1];
|
|
getrandom::getrandom(&mut byte).unwrap_or_default();
|
|
byte[0] % 10
|
|
})
|
|
.collect();
|
|
digits
|
|
.iter()
|
|
.map(|d| char::from_digit(*d as u32, 10).unwrap_or('0'))
|
|
.collect()
|
|
}
|
|
|
|
/// 创建新的短信验证码记录
|
|
///
|
|
/// 同一手机号同一目的 60 秒内不允许重复发送。
|
|
pub async fn create(
|
|
pool: &crate::db::Pool,
|
|
phone: &str,
|
|
code: &str,
|
|
purpose: &str,
|
|
expires_in_secs: u64,
|
|
ip_address: Option<&str>,
|
|
) -> AppResult<SmsCode> {
|
|
let (document_id, now) = id::new_document_id_and_timestamp();
|
|
let expires_at =
|
|
crate::utils::tz::now_utc() + chrono::Duration::seconds(expires_in_secs as i64);
|
|
|
|
let sql = format!(
|
|
"INSERT INTO sms_codes (document_id, phone, code, purpose, expires_at, ip_address, created_at) VALUES ({}, {}, {}, {}, {}, {}, {})",
|
|
ph(1),
|
|
ph(2),
|
|
ph(3),
|
|
ph(4),
|
|
ph(5),
|
|
ph(6),
|
|
ph(7),
|
|
);
|
|
sqlx::query(&sql)
|
|
.bind(&document_id)
|
|
.bind(phone)
|
|
.bind(code)
|
|
.bind(purpose)
|
|
.bind(expires_at)
|
|
.bind(ip_address)
|
|
.bind(now)
|
|
.execute(pool)
|
|
.await?;
|
|
|
|
let sql2 = format!("SELECT * FROM sms_codes WHERE document_id = {}", ph(1));
|
|
sqlx::query_as::<_, SmsCode>(&sql2)
|
|
.bind(&document_id)
|
|
.fetch_optional(pool)
|
|
.await?
|
|
.ok_or_else(|| {
|
|
crate::errors::app_error::AppError::Internal(anyhow::anyhow!(
|
|
"failed to fetch sms code"
|
|
))
|
|
})
|
|
}
|
|
|
|
/// 根据 ID 查找验证码
|
|
pub async fn find_by_id(pool: &crate::db::Pool, id: i64) -> AppResult<Option<SmsCode>> {
|
|
let sql = format!("SELECT * FROM sms_codes WHERE id = {}", ph(1));
|
|
let row = sqlx::query_as::<_, SmsCode>(&sql)
|
|
.bind(id)
|
|
.fetch_optional(pool)
|
|
.await?;
|
|
Ok(row)
|
|
}
|
|
|
|
/// 查找手机号最近的未验证验证码
|
|
pub async fn find_latest_unverified(
|
|
pool: &crate::db::Pool,
|
|
phone: &str,
|
|
purpose: &str,
|
|
) -> AppResult<Option<SmsCode>> {
|
|
let sql = format!(
|
|
"SELECT * FROM sms_codes WHERE phone = {} AND purpose = {} AND verified_at IS NULL ORDER BY created_at DESC LIMIT 1",
|
|
ph(1),
|
|
ph(2),
|
|
);
|
|
let row = sqlx::query_as::<_, SmsCode>(&sql)
|
|
.bind(phone)
|
|
.bind(purpose)
|
|
.fetch_optional(pool)
|
|
.await?;
|
|
Ok(row)
|
|
}
|
|
|
|
/// 检查是否在限流期内(同一手机号同一目的最近 N 秒内是否有发送记录)
|
|
pub async fn is_rate_limited(
|
|
pool: &crate::db::Pool,
|
|
phone: &str,
|
|
purpose: &str,
|
|
within_secs: u64,
|
|
) -> AppResult<bool> {
|
|
let cutoff = crate::utils::tz::now_utc() - chrono::Duration::seconds(within_secs as i64);
|
|
let sql = format!(
|
|
"SELECT COUNT(*) as cnt FROM sms_codes WHERE phone = {} AND purpose = {} AND created_at > {}",
|
|
ph(1),
|
|
ph(2),
|
|
ph(3),
|
|
);
|
|
let row: (i64,) = sqlx::query_as(&sql)
|
|
.bind(phone)
|
|
.bind(purpose)
|
|
.bind(cutoff)
|
|
.fetch_one(pool)
|
|
.await?;
|
|
Ok(row.0 > 0)
|
|
}
|
|
|
|
/// 验证码验证:匹配后标记已验证,错误时增加 attempts
|
|
pub async fn verify_code(
|
|
pool: &crate::db::Pool,
|
|
id: i64,
|
|
input_code: &str,
|
|
) -> AppResult<VerifyResult> {
|
|
let sms = find_by_id(pool, id)
|
|
.await?
|
|
.ok_or_else(|| crate::errors::app_error::AppError::BadRequest("invalid_code".into()))?;
|
|
|
|
if sms.verified_at.is_some() {
|
|
return Ok(VerifyResult::AlreadyUsed);
|
|
}
|
|
|
|
if sms.expires_at < crate::utils::tz::now_utc() {
|
|
return Ok(VerifyResult::Expired);
|
|
}
|
|
|
|
if sms.attempts >= 5 {
|
|
return Ok(VerifyResult::MaxAttempts);
|
|
}
|
|
|
|
if sms.code != input_code {
|
|
let sql = format!(
|
|
"UPDATE sms_codes SET attempts = attempts + 1 WHERE id = {}",
|
|
ph(1),
|
|
);
|
|
sqlx::query(&sql).bind(id).execute(pool).await?;
|
|
return Ok(VerifyResult::WrongCode);
|
|
}
|
|
|
|
let now = crate::utils::tz::now_utc();
|
|
let sql = format!(
|
|
"UPDATE sms_codes SET verified_at = {} WHERE id = {}",
|
|
ph(1),
|
|
ph(2),
|
|
);
|
|
sqlx::query(&sql).bind(now).bind(id).execute(pool).await?;
|
|
|
|
Ok(VerifyResult::Verified)
|
|
}
|
|
|
|
/// 验证结果
|
|
#[derive(Debug, Clone, PartialEq)]
|
|
pub enum VerifyResult {
|
|
Verified,
|
|
WrongCode,
|
|
Expired,
|
|
AlreadyUsed,
|
|
MaxAttempts,
|
|
}
|
|
|
|
/// 清理过期的验证码记录
|
|
pub async fn cleanup_expired(pool: &crate::db::Pool) -> AppResult<u64> {
|
|
let now = crate::utils::tz::now_utc();
|
|
let sql = format!("DELETE FROM sms_codes WHERE expires_at < {}", ph(1));
|
|
let result = sqlx::query(&sql).bind(now).execute(pool).await?;
|
|
Ok(result.rows_affected())
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
async fn setup_pool() -> crate::db::Pool {
|
|
let pool = crate::db::Pool::connect("sqlite::memory:").await.unwrap();
|
|
sqlx::query(crate::db::schema::SCHEMA_SQL)
|
|
.execute(&pool)
|
|
.await
|
|
.unwrap();
|
|
pool
|
|
}
|
|
|
|
fn unique_phone() -> String {
|
|
let id = crate::utils::id::new_document_id();
|
|
let hash = id
|
|
.bytes()
|
|
.fold(0u32, |acc, b| acc.wrapping_mul(31).wrapping_add(b as u32));
|
|
format!("1380000{:04}", hash % 10000)
|
|
}
|
|
|
|
#[test]
|
|
fn generate_code_length() {
|
|
let code = super::generate_code(6);
|
|
assert_eq!(code.len(), 6);
|
|
assert!(code.chars().all(|c| c.is_ascii_digit()));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn create_and_find_by_id() {
|
|
let pool = setup_pool().await;
|
|
let phone = unique_phone();
|
|
let code = super::generate_code(6);
|
|
|
|
let sms = super::create(&pool, &phone, &code, "login", 300, Some("127.0.0.1"))
|
|
.await
|
|
.unwrap();
|
|
|
|
let found = super::find_by_id(&pool, sms.id).await.unwrap();
|
|
assert!(found.is_some());
|
|
let row = found.unwrap();
|
|
assert_eq!(row.phone, phone);
|
|
assert_eq!(row.code, code);
|
|
assert_eq!(row.purpose, "login");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn find_latest_unverified() {
|
|
let pool = setup_pool().await;
|
|
let phone = unique_phone();
|
|
|
|
let _first = super::create(&pool, &phone, "111111", "login", 300, None)
|
|
.await
|
|
.unwrap();
|
|
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
|
|
let second = super::create(&pool, &phone, "222222", "login", 300, None)
|
|
.await
|
|
.unwrap();
|
|
|
|
let latest = super::find_latest_unverified(&pool, &phone, "login")
|
|
.await
|
|
.unwrap();
|
|
assert!(latest.is_some());
|
|
assert_eq!(latest.unwrap().id, second.id);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn verify_code_correct() {
|
|
let pool = setup_pool().await;
|
|
let phone = unique_phone();
|
|
let code = "654321";
|
|
|
|
let sms = super::create(&pool, &phone, code, "login", 300, None)
|
|
.await
|
|
.unwrap();
|
|
|
|
let result = super::verify_code(&pool, sms.id, code).await.unwrap();
|
|
assert_eq!(result, super::VerifyResult::Verified);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn verify_code_wrong() {
|
|
let pool = setup_pool().await;
|
|
let phone = unique_phone();
|
|
|
|
let sms = super::create(&pool, &phone, "123456", "login", 300, None)
|
|
.await
|
|
.unwrap();
|
|
|
|
let result = super::verify_code(&pool, sms.id, "000000").await.unwrap();
|
|
assert_eq!(result, super::VerifyResult::WrongCode);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn is_rate_limited() {
|
|
let pool = setup_pool().await;
|
|
let phone = unique_phone();
|
|
|
|
super::create(&pool, &phone, "111111", "login", 300, None)
|
|
.await
|
|
.unwrap();
|
|
|
|
let limited = super::is_rate_limited(&pool, &phone, "login", 60)
|
|
.await
|
|
.unwrap();
|
|
assert!(limited);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn cleanup_expired() {
|
|
let pool = setup_pool().await;
|
|
let phone = unique_phone();
|
|
|
|
super::create(&pool, &phone, "111111", "login", 0, None)
|
|
.await
|
|
.unwrap();
|
|
|
|
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
|
|
|
|
let removed = super::cleanup_expired(&pool).await.unwrap();
|
|
assert!(removed > 0);
|
|
}
|
|
}
|