更新
This commit is contained in:
@@ -0,0 +1,78 @@
|
||||
"""邮箱验证令牌创建与校验。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import secrets
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from sqlalchemy import delete, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from models.models import EmailVerificationToken, User
|
||||
|
||||
|
||||
def _token_expires_at(hours: int) -> datetime:
|
||||
return datetime.utcnow() + timedelta(hours=max(1, min(168, int(hours or 24))))
|
||||
|
||||
|
||||
async def invalidate_user_tokens(db: AsyncSession, user_id: int) -> None:
|
||||
await db.execute(
|
||||
delete(EmailVerificationToken).where(EmailVerificationToken.user_id == user_id)
|
||||
)
|
||||
|
||||
|
||||
async def create_verification_token(
|
||||
db: AsyncSession,
|
||||
user: User,
|
||||
expire_hours: int = 24,
|
||||
) -> str:
|
||||
await invalidate_user_tokens(db, user.id)
|
||||
token = secrets.token_urlsafe(32)
|
||||
db.add(
|
||||
EmailVerificationToken(
|
||||
user_id=user.id,
|
||||
token=token,
|
||||
expires_at=_token_expires_at(expire_hours),
|
||||
)
|
||||
)
|
||||
await db.flush()
|
||||
return token
|
||||
|
||||
|
||||
async def verify_email_token(db: AsyncSession, token: str) -> User | None:
|
||||
value = (token or "").strip()
|
||||
if not value:
|
||||
return None
|
||||
|
||||
result = await db.execute(
|
||||
select(EmailVerificationToken, User)
|
||||
.join(User, User.id == EmailVerificationToken.user_id)
|
||||
.where(EmailVerificationToken.token == value)
|
||||
)
|
||||
row = result.first()
|
||||
if not row:
|
||||
return None
|
||||
|
||||
record, user = row
|
||||
if record.expires_at < datetime.utcnow():
|
||||
await db.delete(record)
|
||||
await db.flush()
|
||||
return None
|
||||
|
||||
user.email_verified = True
|
||||
user.email_verified_at = datetime.utcnow()
|
||||
await db.delete(record)
|
||||
await db.flush()
|
||||
return user
|
||||
|
||||
|
||||
def mask_email(email: str) -> str:
|
||||
value = (email or "").strip()
|
||||
if "@" not in value:
|
||||
return value
|
||||
local, domain = value.split("@", 1)
|
||||
if len(local) <= 2:
|
||||
masked_local = local[0] + "*"
|
||||
else:
|
||||
masked_local = local[0] + "*" * (len(local) - 2) + local[-1]
|
||||
return f"{masked_local}@{domain}"
|
||||
Reference in New Issue
Block a user