Files
video-gen/video-gen-api/app/api/v1/auth.py
T
2026-06-16 13:12:43 +08:00

316 lines
10 KiB
Python

from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.dependencies import (
get_current_user,
get_current_user_allow_password_pending,
get_db,
)
from app.models.system_config import SystemConfig
from app.models.user import User
from app.schemas.auth import (
ChangePasswordRequest,
LoginRequest,
RegisterRequest,
SetPasswordRequest,
SmsLoginRequest,
)
from app.schemas.user import UserOut
from app.services.auth import (
authenticate_user,
create_access_token,
decode_access_token,
get_user_by_phone,
hash_password,
verify_password,
)
from app.services.sms import verify_sms_code
from app.utils.id_gen import generate_id
router = APIRouter(prefix="/auth", tags=["auth"])
def _validate_captcha_if_needed(captcha_token: str | None) -> None:
# 保留原密码登录的图形验证码逻辑,不改成短信验证码。
if settings.SMS_MOCK:
return
if not captcha_token:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="需要验证码",
)
token_sub = decode_access_token(captcha_token)
if not token_sub or not token_sub.startswith("captcha:"):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="验证码无效或已过期",
)
def _token_response(user: User, remember_me: bool = False) -> dict:
user.credits = round(user.credits, 2)
token = create_access_token(user.id, remember_me)
return {
"access_token": token,
"token_type": "bearer",
"user": UserOut.model_validate(user),
"must_set_password": bool(user.user_type == "frontend" and not user.hashed_password),
}
async def _assign_default_frontend_menus(db: AsyncSession, user: User) -> None:
from app.models.menu_config import MenuConfig
result = await db.execute(
select(MenuConfig).where(
MenuConfig.is_default == True,
MenuConfig.is_active == True,
MenuConfig.menu_target.in_(["frontend", "both"]),
MenuConfig.menu_type == "page",
)
)
default_menus = result.scalars().all()
if default_menus:
user.allowed_menus = [m.path for m in default_menus if m.path]
async def _get_register_credits(db: AsyncSession) -> int:
result = await db.execute(
select(SystemConfig.value).where(SystemConfig.key == "user_register_credits").limit(1)
)
value = result.scalar_one_or_none()
return int(value) if value else 100
async def _handle_daily_login_credits(db: AsyncSession, user: User) -> None:
enabled_result = await db.execute(
select(SystemConfig.value).where(SystemConfig.key == "user_login_credits_enabled").limit(1)
)
enabled = enabled_result.scalar_one_or_none() == "true"
if not enabled:
return
credits_result = await db.execute(
select(SystemConfig.value).where(SystemConfig.key == "user_login_credits").limit(1)
)
credits = int(credits_result.scalar_one_or_none() or "0")
if credits <= 0:
return
today = datetime.now().date()
if user.last_login_at:
last_login_date = user.last_login_at.date()
if last_login_date >= today:
return
user.credits += credits
@router.post(
"/login",
summary="客户端密码登录",
description="保留原有用户名/手机号 + 密码登录。仅允许 frontend 用户登录;管理员仍使用 /auth/admin-login。",
)
async def login(req: LoginRequest, db: AsyncSession = Depends(get_db)):
_validate_captcha_if_needed(req.captcha_token)
user = await authenticate_user(db, req.username, req.password)
if not user:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户名或密码错误",
)
# Only allow frontend users to login via this endpoint
if user.user_type != "frontend":
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="该账号不允许在此登录",
)
await _handle_daily_login_credits(db, user)
user.last_login_at = datetime.now()
await db.flush()
return _token_response(user, req.remember_me)
@router.post(
"/sms-login",
summary="客户端短信验证码登录",
description="新增兼容登录方式:手机号 + 短信验证码登录。不覆盖 /auth/login 密码登录。仅允许 frontend 用户登录。",
)
async def sms_login(req: SmsLoginRequest, db: AsyncSession = Depends(get_db)):
ok = await verify_sms_code(req.phone, req.code, "login")
if not ok:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="验证码错误或已过期",
)
user = await get_user_by_phone(db, req.phone)
if not user or not user.is_active:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="账号不存在或已禁用",
)
if user.user_type != "frontend":
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="该账号不允许在此登录",
)
await _handle_daily_login_credits(db, user)
user.last_login_at = datetime.now()
await db.flush()
return _token_response(user, req.remember_me)
@router.post(
"/register",
summary="客户端手机号短信注册",
description="手机号 + 注册短信验证码注册。注册成功后 username 默认等于手机号,不生成密码;前端需根据 must_set_password 引导用户设置密码。",
)
async def register(req: RegisterRequest, db: AsyncSession = Depends(get_db)):
ok = await verify_sms_code(req.phone, req.code, "register")
if not ok:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="验证码错误或已过期",
)
existing_phone = await db.execute(select(User).where(User.phone == req.phone).limit(1))
if existing_phone.scalar_one_or_none():
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="该手机号已注册",
)
register_credits = await _get_register_credits(db)
user = User(
id=generate_id(),
username=req.phone,
phone=req.phone,
hashed_password=hash_password(req.password),
password_set_at=datetime.now(),
credits=register_credits,
is_admin=False,
user_type="frontend",
)
db.add(user)
await db.flush()
await _assign_default_frontend_menus(db, user)
user.credits = round(user.credits, 2)
token = create_access_token(user.id)
return {
"access_token": token,
"token_type": "bearer",
"user": UserOut.model_validate(user),
"must_set_password": True,
}
@router.post("/logout")
async def logout(current_user: User = Depends(get_current_user_allow_password_pending)):
return {"message": "ok"}
@router.get("/me", response_model=UserOut)
async def get_me(current_user: User = Depends(get_current_user_allow_password_pending)):
current_user.credits = round(current_user.credits, 2)
return current_user
@router.post(
"/set-password",
summary="设置登录密码",
description="短信注册或短信登录后,用户没有密码时调用该接口设置密码。该接口允许未设置密码用户访问。",
)
async def set_password(
req: SetPasswordRequest,
current_user: User = Depends(get_current_user_allow_password_pending),
db: AsyncSession = Depends(get_db),
):
if len(req.new_password) < 6:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="密码至少6位",
)
current_user.hashed_password = hash_password(req.new_password)
current_user.password_set_at = datetime.now()
await db.flush()
return {"message": "密码设置成功", "must_set_password": False}
@router.post("/change-password")
async def change_password(
req: ChangePasswordRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
if not verify_password(req.old_password, current_user.hashed_password):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="原密码错误",
)
if len(req.new_password) < 6:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="密码至少6位",
)
current_user.hashed_password = hash_password(req.new_password)
current_user.password_set_at = datetime.now()
await db.flush()
return {"message": "密码修改成功"}
@router.get("/site-info")
async def get_site_info(db: AsyncSession = Depends(get_db)):
"""Public endpoint returning site name and logo."""
result = await db.execute(
select(SystemConfig).where(SystemConfig.key.in_([
"site_name", "site_logo", "user_agreement_url", "privacy_policy_url"
]))
)
configs = result.scalars().all()
info = {c.key: c.value for c in configs}
return {
"site_name": info.get("site_name", "VideoGen.AI"),
"site_logo": info.get("site_logo", ""),
"user_agreement_url": info.get("user_agreement_url", ""),
"privacy_policy_url": info.get("privacy_policy_url", ""),
}
@router.post("/admin-login")
async def admin_login(req: LoginRequest, db: AsyncSession = Depends(get_db)):
"""Admin-only login endpoint."""
user = await authenticate_user(db, req.username, req.password)
if not user:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户名或密码错误",
)
if user.user_type != "admin":
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="该账号不是管理员账号",
)
user.last_login_at = datetime.now()
await db.flush()
token = create_access_token(user.id, req.remember_me)
user.credits = round(user.credits, 2)
return {"access_token": token, "token_type": "bearer", "user": UserOut.model_validate(user)}