130 lines
3.9 KiB
Python
130 lines
3.9 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
认证状态
|
||
"""
|
||
import reflex as rx
|
||
|
||
import re
|
||
from sqlmodel import select
|
||
from application.models.user import EmailVerifyCode
|
||
from application.database.base import get_async_db
|
||
|
||
|
||
class AuthState(rx.State):
|
||
"""
|
||
认证状态
|
||
"""
|
||
|
||
# 当前登录的用户唯一标识(通过本地存储同步)
|
||
user_id: str = rx.LocalStorage("user_id", sync=True)
|
||
|
||
email: str = ""
|
||
captcha: str = ""
|
||
is_policies_agreed: bool = False
|
||
|
||
@rx.event
|
||
async def check(self) -> None:
|
||
"""
|
||
检查认证状态
|
||
"""
|
||
|
||
self.user_id = ""
|
||
|
||
@rx.event
|
||
def set_email(self, email: str) -> None:
|
||
"""
|
||
设置邮箱
|
||
:param email: 邮箱
|
||
:return: None
|
||
"""
|
||
self.email = email.strip()
|
||
|
||
@rx.event
|
||
def set_captcha(self, captcha: str) -> None:
|
||
"""
|
||
设置验证码
|
||
:param captcha: 验证码
|
||
:return: None
|
||
"""
|
||
self.captcha = captcha.strip()
|
||
|
||
@rx.event
|
||
def toggle_policies_agreed(self) -> None:
|
||
"""
|
||
切换协议同意状态
|
||
:return: None
|
||
"""
|
||
self.is_policies_agreed = not self.is_policies_agreed
|
||
|
||
@rx.event
|
||
async def send_email_code(self):
|
||
"""生成并存储邮箱验证码(异步数据库,不阻塞事件循环)"""
|
||
|
||
email_rule = r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$"
|
||
if not re.fullmatch(email_rule, self.email):
|
||
return rx.toast("请输入合法邮箱地址", color="red")
|
||
|
||
# 异步数据库会话
|
||
async for db in get_async_db():
|
||
# 清理该邮箱旧验证码(await 异步查询)
|
||
old_codes = await db.exec(
|
||
select(EmailVerifyCode).where(EmailVerifyCode.email == self.email)
|
||
)
|
||
old_list = old_codes.all()
|
||
for item in old_list:
|
||
await db.delete(item)
|
||
|
||
# 6位验证码,5分钟有效期
|
||
import random
|
||
from datetime import datetime, timedelta
|
||
|
||
code_str = "".join(random.choices("0123456789", k=6))
|
||
expire_time = datetime.utcnow() + timedelta(minutes=5)
|
||
db.add(
|
||
EmailVerifyCode(email=self.email, code=code_str, expire_at=expire_time)
|
||
)
|
||
await db.commit()
|
||
|
||
# 本地调试打印验证码
|
||
print(f"邮箱 {self.email} 测试验证码:{code_str}")
|
||
return rx.toast("验证码已发送(控制台查看测试码)", color="green")
|
||
|
||
@rx.event
|
||
async def login_by_email_code(self):
|
||
"""邮箱验证码登录,不存在用户自动注册"""
|
||
if not self.email:
|
||
return rx.toast("请填写邮箱", color="red")
|
||
if not self.code:
|
||
return rx.toast("请填写验证码", color="red")
|
||
if not self.agree_policy:
|
||
return rx.toast("请勾选同意用户协议与隐私政策", color="red")
|
||
|
||
db = get_db_session()
|
||
# 查询有效验证码
|
||
code_record = db.exec(
|
||
select(EmailVerifyCode).where(
|
||
EmailVerifyCode.email == self.email, EmailVerifyCode.code == self.code
|
||
)
|
||
).first()
|
||
if not code_record or code_record.is_expired():
|
||
return rx.toast("验证码错误或已过期", color="red")
|
||
|
||
# 一次性验证码使用后删除
|
||
db.delete(code_record)
|
||
|
||
# 查找用户,无则新建
|
||
user = db.exec(select(User).where(User.email == self.email)).first()
|
||
if not user:
|
||
user = User(email=self.email, nickname=self.email.split("@")[0])
|
||
db.add(user)
|
||
user.last_login_at = datetime.utcnow()
|
||
db.commit()
|
||
db.refresh(user)
|
||
|
||
# 登录成功,弹窗自动消失
|
||
self.user_id = str(user.id)
|
||
self.email = ""
|
||
self.code = ""
|
||
self.agree_policy = False
|
||
return rx.toast("登录成功", color="green")
|