- models: 删除 Profile 模型,字段合并到 User;User 新增云朵/公开/收藏计数器;Cloud 新增 favorite_count - schemas/serializers: 展平 UserOut/AuthOut/MeOut/AdminUserOut,移除嵌套 profile - deps: user.profile.role→user.role,移除 selectinload profile - auth/register: 注册时直接设置 user 字段,不再创建 Profile - clouds: 创建/删除云朵时同步 user.cloud_count 计数器,状态变化时同步 public_cloud_count - favorites: 点赞/取消时同步 cloud.favorite_count 计数器 - admin: 审批/隐藏/批量删时同步 public_cloud_count 计数器 - profiles/stats: 用计数器替代实时 COUNT 查询 - alembic: 新增 migration 合并 profiles 到 users,初始化 counters
355 lines
12 KiB
Python
355 lines
12 KiB
Python
import logging
|
|
from datetime import datetime, timedelta, timezone
|
|
|
|
from fastapi import APIRouter, HTTPException, Request, Response, status
|
|
from sqlalchemy import select, update
|
|
from sqlalchemy.exc import IntegrityError
|
|
from sqlalchemy.orm import selectinload
|
|
|
|
from app.config import settings
|
|
from app.deps import CurrentUser, DbSession
|
|
from app.models import EmailToken, RefreshSession, User
|
|
from app.schemas import (
|
|
AuthOut,
|
|
ChangePasswordIn,
|
|
ForgotPasswordIn,
|
|
LoginIn,
|
|
MeOut,
|
|
MessageOut,
|
|
RegisterIn,
|
|
ResetPasswordIn,
|
|
TokenIn,
|
|
)
|
|
from app.security import (
|
|
DUMMY_PASSWORD_HASH,
|
|
create_access_token,
|
|
create_random_token,
|
|
hash_password,
|
|
hash_token,
|
|
refresh_expires_at,
|
|
utc_now,
|
|
verify_password,
|
|
)
|
|
from app.serializers import user_out
|
|
from app.services.email import send_confirmation_email, send_password_reset_email
|
|
|
|
|
|
router = APIRouter(prefix="/auth", tags=["认证"])
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _aware_utc(value: datetime) -> datetime:
|
|
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value.astimezone(timezone.utc)
|
|
|
|
|
|
def _client_metadata(request: Request) -> tuple[str | None, str | None]:
|
|
user_agent = request.headers.get("user-agent")
|
|
ip_address = request.client.host if request.client else None
|
|
return user_agent, ip_address
|
|
|
|
|
|
def _set_refresh_cookie(response: Response, token: str) -> None:
|
|
response.set_cookie(
|
|
key=settings.refresh_cookie_name,
|
|
value=token,
|
|
max_age=settings.refresh_token_expire_days * 24 * 60 * 60,
|
|
httponly=True,
|
|
secure=settings.cookie_secure,
|
|
samesite=settings.cookie_samesite,
|
|
domain=settings.cookie_domain,
|
|
path=f"{settings.api_v1_prefix}/auth",
|
|
)
|
|
|
|
|
|
def _clear_refresh_cookie(response: Response) -> None:
|
|
response.delete_cookie(
|
|
key=settings.refresh_cookie_name,
|
|
domain=settings.cookie_domain,
|
|
path=f"{settings.api_v1_prefix}/auth",
|
|
secure=settings.cookie_secure,
|
|
httponly=True,
|
|
samesite=settings.cookie_samesite,
|
|
)
|
|
|
|
|
|
async def _create_email_token(db: DbSession, user: User, purpose: str, expires_delta: timedelta) -> str:
|
|
now = utc_now()
|
|
await db.execute(
|
|
update(EmailToken)
|
|
.where(
|
|
EmailToken.user_id == user.id,
|
|
EmailToken.purpose == purpose,
|
|
EmailToken.used_at.is_(None),
|
|
)
|
|
.values(used_at=now)
|
|
)
|
|
token = create_random_token()
|
|
db.add(
|
|
EmailToken(
|
|
user_id=user.id,
|
|
purpose=purpose,
|
|
token_hash=hash_token(token),
|
|
expires_at=now + expires_delta,
|
|
)
|
|
)
|
|
return token
|
|
|
|
|
|
@router.post("/register", response_model=MessageOut, status_code=status.HTTP_201_CREATED)
|
|
async def register(payload: RegisterIn, db: DbSession) -> MessageOut:
|
|
email = str(payload.email).strip().lower()
|
|
user = User(email=email, password_hash=hash_password(payload.password), username=payload.username)
|
|
db.add(user)
|
|
try:
|
|
await db.flush()
|
|
token = await _create_email_token(
|
|
db,
|
|
user,
|
|
"confirm_email",
|
|
timedelta(hours=settings.email_confirmation_expire_hours),
|
|
)
|
|
await db.commit()
|
|
except IntegrityError as exc:
|
|
await db.rollback()
|
|
message = str(exc.orig).lower()
|
|
if "email" in message:
|
|
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="该邮箱已被注册") from exc
|
|
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="这个昵称已经被使用") from exc
|
|
|
|
try:
|
|
await send_confirmation_email(user.email, token)
|
|
except Exception:
|
|
logger.exception("注册成功,但确认邮件发送失败:%s", user.email)
|
|
return MessageOut(message="注册成功,请查收确认邮件")
|
|
|
|
|
|
@router.post("/resend-confirmation", response_model=MessageOut, status_code=status.HTTP_202_ACCEPTED)
|
|
async def resend_confirmation(payload: ForgotPasswordIn, db: DbSession) -> MessageOut:
|
|
email = str(payload.email).strip().lower()
|
|
result = await db.execute(select(User).where(User.email == email))
|
|
user = result.scalar_one_or_none()
|
|
if user and user.email_verified_at is None:
|
|
token = await _create_email_token(
|
|
db,
|
|
user,
|
|
"confirm_email",
|
|
timedelta(hours=settings.email_confirmation_expire_hours),
|
|
)
|
|
await db.commit()
|
|
try:
|
|
await send_confirmation_email(user.email, token)
|
|
except Exception:
|
|
logger.exception("确认邮件重发失败:%s", user.email)
|
|
return MessageOut(message="如果账号存在且尚未确认,确认邮件将会发送")
|
|
|
|
|
|
@router.post("/confirm-email", response_model=MessageOut)
|
|
async def confirm_email(payload: TokenIn, db: DbSession) -> MessageOut:
|
|
now = utc_now()
|
|
result = await db.execute(
|
|
select(EmailToken)
|
|
.where(
|
|
EmailToken.token_hash == hash_token(payload.token),
|
|
EmailToken.purpose == "confirm_email",
|
|
EmailToken.used_at.is_(None),
|
|
EmailToken.expires_at > now,
|
|
)
|
|
.with_for_update()
|
|
)
|
|
email_token = result.scalar_one_or_none()
|
|
if not email_token:
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="确认链接无效或已过期")
|
|
user = await db.get(User, email_token.user_id)
|
|
if not user:
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="账号不存在")
|
|
user.email_verified_at = user.email_verified_at or now
|
|
email_token.used_at = now
|
|
await db.commit()
|
|
return MessageOut(message="邮箱确认成功")
|
|
|
|
|
|
@router.post("/login", response_model=AuthOut)
|
|
async def login(payload: LoginIn, request: Request, response: Response, db: DbSession) -> AuthOut:
|
|
email = str(payload.email).strip().lower()
|
|
result = await db.execute(select(User).where(User.email == email))
|
|
user = result.scalar_one_or_none()
|
|
encoded_hash = user.password_hash if user else DUMMY_PASSWORD_HASH
|
|
password_valid = verify_password(payload.password, encoded_hash)
|
|
if not user or not password_valid:
|
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="邮箱或密码错误")
|
|
if user.email_verified_at is None:
|
|
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="邮箱尚未确认")
|
|
if user.is_disabled:
|
|
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="账号已被禁用")
|
|
|
|
token = create_random_token()
|
|
user_agent, ip_address = _client_metadata(request)
|
|
session = RefreshSession(
|
|
user_id=user.id,
|
|
token_hash=hash_token(token),
|
|
expires_at=refresh_expires_at(),
|
|
user_agent=user_agent,
|
|
ip_address=ip_address,
|
|
)
|
|
db.add(session)
|
|
await db.commit()
|
|
_set_refresh_cookie(response, token)
|
|
return AuthOut(
|
|
access_token=create_access_token(user.id, session.id),
|
|
expires_in=settings.access_token_expire_minutes * 60,
|
|
user=user_out(user),
|
|
)
|
|
|
|
|
|
@router.post("/refresh", response_model=AuthOut)
|
|
async def refresh(request: Request, response: Response, db: DbSession) -> AuthOut:
|
|
raw_token = request.cookies.get(settings.refresh_cookie_name)
|
|
if not raw_token:
|
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="缺少刷新令牌")
|
|
now = utc_now()
|
|
result = await db.execute(
|
|
select(RefreshSession)
|
|
.options(selectinload(RefreshSession.user))
|
|
.where(RefreshSession.token_hash == hash_token(raw_token))
|
|
.with_for_update()
|
|
)
|
|
old_session = result.scalar_one_or_none()
|
|
if (
|
|
not old_session
|
|
or old_session.revoked_at is not None
|
|
or _aware_utc(old_session.expires_at) <= now
|
|
):
|
|
_clear_refresh_cookie(response)
|
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="刷新令牌无效或已过期")
|
|
user = old_session.user
|
|
if user.is_disabled:
|
|
old_session.revoked_at = now
|
|
await db.commit()
|
|
_clear_refresh_cookie(response)
|
|
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="账号已被禁用")
|
|
|
|
token = create_random_token()
|
|
user_agent, ip_address = _client_metadata(request)
|
|
new_session = RefreshSession(
|
|
user_id=user.id,
|
|
token_hash=hash_token(token),
|
|
expires_at=refresh_expires_at(),
|
|
user_agent=user_agent,
|
|
ip_address=ip_address,
|
|
)
|
|
db.add(new_session)
|
|
await db.flush()
|
|
old_session.revoked_at = now
|
|
old_session.replaced_by_id = new_session.id
|
|
await db.commit()
|
|
_set_refresh_cookie(response, token)
|
|
return AuthOut(
|
|
access_token=create_access_token(user.id, new_session.id),
|
|
expires_in=settings.access_token_expire_minutes * 60,
|
|
user=user_out(user),
|
|
)
|
|
|
|
|
|
@router.post("/logout", response_model=MessageOut)
|
|
async def logout(request: Request, response: Response, db: DbSession) -> MessageOut:
|
|
raw_token = request.cookies.get(settings.refresh_cookie_name)
|
|
if raw_token:
|
|
result = await db.execute(
|
|
select(RefreshSession).where(RefreshSession.token_hash == hash_token(raw_token))
|
|
)
|
|
session = result.scalar_one_or_none()
|
|
if session and session.revoked_at is None:
|
|
session.revoked_at = utc_now()
|
|
await db.commit()
|
|
_clear_refresh_cookie(response)
|
|
return MessageOut(message="已退出登录")
|
|
|
|
|
|
@router.get("/me", response_model=MeOut)
|
|
async def me(user: CurrentUser) -> MeOut:
|
|
return MeOut(user=user_out(user))
|
|
|
|
|
|
@router.post("/forgot-password", response_model=MessageOut, status_code=status.HTTP_202_ACCEPTED)
|
|
async def forgot_password(payload: ForgotPasswordIn, db: DbSession) -> MessageOut:
|
|
email = str(payload.email).strip().lower()
|
|
result = await db.execute(select(User).where(User.email == email))
|
|
user = result.scalar_one_or_none()
|
|
if user and user.email_verified_at is not None:
|
|
token = await _create_email_token(
|
|
db,
|
|
user,
|
|
"reset_password",
|
|
timedelta(minutes=settings.password_reset_expire_minutes),
|
|
)
|
|
await db.commit()
|
|
try:
|
|
await send_password_reset_email(user.email, token)
|
|
except Exception:
|
|
logger.exception("密码重置邮件发送失败:%s", user.email)
|
|
return MessageOut(message="如果该邮箱已经注册,密码重置邮件将会发送")
|
|
|
|
|
|
@router.post("/reset-password", response_model=MessageOut)
|
|
async def reset_password(payload: ResetPasswordIn, db: DbSession) -> MessageOut:
|
|
now = utc_now()
|
|
result = await db.execute(
|
|
select(EmailToken)
|
|
.where(
|
|
EmailToken.token_hash == hash_token(payload.token),
|
|
EmailToken.purpose == "reset_password",
|
|
EmailToken.used_at.is_(None),
|
|
EmailToken.expires_at > now,
|
|
)
|
|
.with_for_update()
|
|
)
|
|
email_token = result.scalar_one_or_none()
|
|
if not email_token:
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="重置链接无效或已过期")
|
|
user = await db.get(User, email_token.user_id)
|
|
if not user:
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="账号不存在")
|
|
user.password_hash = hash_password(payload.password)
|
|
email_token.used_at = now
|
|
await db.execute(
|
|
update(RefreshSession)
|
|
.where(RefreshSession.user_id == user.id, RefreshSession.revoked_at.is_(None))
|
|
.values(revoked_at=now)
|
|
)
|
|
await db.commit()
|
|
return MessageOut(message="密码已重置,请重新登录")
|
|
|
|
|
|
@router.patch("/password", response_model=AuthOut)
|
|
async def change_password(
|
|
payload: ChangePasswordIn,
|
|
request: Request,
|
|
response: Response,
|
|
user: CurrentUser,
|
|
db: DbSession,
|
|
) -> AuthOut:
|
|
now = utc_now()
|
|
user.password_hash = hash_password(payload.password)
|
|
await db.execute(
|
|
update(RefreshSession)
|
|
.where(RefreshSession.user_id == user.id, RefreshSession.revoked_at.is_(None))
|
|
.values(revoked_at=now)
|
|
)
|
|
token = create_random_token()
|
|
user_agent, ip_address = _client_metadata(request)
|
|
new_session = RefreshSession(
|
|
user_id=user.id,
|
|
token_hash=hash_token(token),
|
|
expires_at=refresh_expires_at(),
|
|
user_agent=user_agent,
|
|
ip_address=ip_address,
|
|
)
|
|
db.add(new_session)
|
|
await db.commit()
|
|
_set_refresh_cookie(response, token)
|
|
return AuthOut(
|
|
access_token=create_access_token(user.id, new_session.id),
|
|
expires_in=settings.access_token_expire_minutes * 60,
|
|
user=user_out(user),
|
|
)
|