3
This commit is contained in:
@@ -0,0 +1,86 @@
|
||||
"""认证服务:注册、登录、登出、Session 校验。"""
|
||||
|
||||
from app.core.errors import (
|
||||
AuthRequiredError,
|
||||
ConflictError,
|
||||
InvalidCredentialsError,
|
||||
)
|
||||
from app.core.security import hash_password, verify_password
|
||||
from app.core.session import create_session, delete_session, get_session_user
|
||||
from app.models.user import User
|
||||
from app.repositories.plan_repo import PlanRepository
|
||||
from app.repositories.user_repo import UserRepository
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
class AuthService:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
self._user_repo = UserRepository(session)
|
||||
self._plan_repo = PlanRepository(session)
|
||||
|
||||
def register(self, username: str, email: str, password: str) -> tuple[User, str]:
|
||||
"""注册新用户。返回 (user, session_token)。
|
||||
|
||||
Raises:
|
||||
ConflictError: 用户名或邮箱已存在
|
||||
"""
|
||||
username_taken, email_taken = self._user_repo.exists_username_or_email(username, email)
|
||||
if username_taken:
|
||||
raise ConflictError("用户名已被占用。", code="USERNAME_TAKEN")
|
||||
if email_taken:
|
||||
raise ConflictError("邮箱已被注册。", code="EMAIL_TAKEN")
|
||||
|
||||
plan = self._plan_repo.get_or_create_free()
|
||||
user = self._user_repo.create(
|
||||
username=username,
|
||||
email=email,
|
||||
password_hash=hash_password(password),
|
||||
plan_id=plan.id,
|
||||
)
|
||||
self._session.commit()
|
||||
|
||||
token = create_session(user.id)
|
||||
return user, token
|
||||
|
||||
def login(self, username_or_email: str, password: str) -> tuple[User, str]:
|
||||
"""登录。返回 (user, session_token)。
|
||||
|
||||
Raises:
|
||||
InvalidCredentialsError: 用户名/邮箱或密码错误
|
||||
"""
|
||||
user = self._user_repo.get_by_username_or_email(username_or_email)
|
||||
if user is None:
|
||||
raise InvalidCredentialsError()
|
||||
if not verify_password(password, user.password_hash):
|
||||
raise InvalidCredentialsError()
|
||||
if user.status != "active":
|
||||
raise InvalidCredentialsError("账户已被禁用。")
|
||||
|
||||
token = create_session(user.id)
|
||||
return user, token
|
||||
|
||||
def logout(self, session_token: str) -> None:
|
||||
"""登出:删除 session。"""
|
||||
delete_session(session_token)
|
||||
|
||||
def get_current_user(self, session_token: str | None) -> User:
|
||||
"""根据 session token 获取当前用户。
|
||||
|
||||
Raises:
|
||||
AuthRequiredError: 未登录或 session 已过期
|
||||
"""
|
||||
if not session_token:
|
||||
raise AuthRequiredError()
|
||||
user_id = get_session_user(session_token)
|
||||
if user_id is None:
|
||||
raise AuthRequiredError("登录已过期,请重新登录。")
|
||||
user = self._user_repo.get_by_id(user_id)
|
||||
if user is None or user.status != "active":
|
||||
raise AuthRequiredError()
|
||||
return user
|
||||
|
||||
def change_password(self, user: User, new_password: str) -> None:
|
||||
"""修改密码。"""
|
||||
self._user_repo.update_password(user, hash_password(new_password))
|
||||
self._session.commit()
|
||||
Reference in New Issue
Block a user