109 lines
4.0 KiB
Python
109 lines
4.0 KiB
Python
"""认证服务:注册、登录、登出、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, role: str = "customer") -> tuple[User, str]:
|
|
"""注册新用户。返回 (user, session_token)。
|
|
|
|
Args:
|
|
role: "customer"(默认)或 "internal"
|
|
|
|
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,
|
|
role=role,
|
|
)
|
|
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 internal_login(self, username_or_email: str, password: str) -> tuple[User, str]:
|
|
"""内部登录(仅 internal 角色可用)。返回 (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("账户已被禁用。")
|
|
if user.role != "internal":
|
|
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() |