362 lines
12 KiB
Python
362 lines
12 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.credit_record import CreditRecord
|
|
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.services.resource_capacity_service import get_user_resource_capacity_usage
|
|
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 _add_register_credit_record(db: AsyncSession, user: User, credits: int) -> None:
|
|
if credits <= 0:
|
|
return
|
|
record = CreditRecord(
|
|
id=generate_id(),
|
|
user_id=user.id,
|
|
type="recharge",
|
|
amount=credits,
|
|
balance_after=user.credits,
|
|
description=f"注册赠送 {credits} 积分",
|
|
)
|
|
db.add(record)
|
|
|
|
|
|
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
|
|
|
|
record = CreditRecord(
|
|
id=generate_id(),
|
|
user_id=user.id,
|
|
type="recharge",
|
|
amount=credits,
|
|
balance_after=user.credits,
|
|
description=f"每日登录赠送 {credits} 积分",
|
|
)
|
|
db.add(record)
|
|
|
|
|
|
@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 _add_register_credit_record(db, user, register_credits)
|
|
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),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
current_user.username = "用户"+current_user.username[-4:] if current_user.username == current_user.phone else current_user.username
|
|
current_user.credits = round(current_user.credits, 2)
|
|
resource_capacity = await get_user_resource_capacity_usage(db, current_user.id)
|
|
return UserOut.model_validate(current_user).model_copy(
|
|
update={"resource_capacity": resource_capacity}
|
|
)
|
|
|
|
|
|
@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, logo, agreement and copyright info."""
|
|
result = await db.execute(
|
|
select(SystemConfig).where(SystemConfig.key.in_([
|
|
"site_name", "site_logo", "user_agreement_privacy_url", "site_copyright"
|
|
]))
|
|
)
|
|
configs = result.scalars().all()
|
|
info = {c.key: c.value for c in configs}
|
|
|
|
# Helper function to convert relative path to full URL
|
|
def to_full_url(path: str | None) -> str:
|
|
if not path:
|
|
return ""
|
|
# If already absolute URL, return as-is
|
|
if path.startswith("http://") or path.startswith("https://"):
|
|
return path
|
|
# Convert relative path to full URL
|
|
base_url = settings.BASE_URL.rstrip("/")
|
|
return f"{base_url}{path}"
|
|
|
|
return {
|
|
"site_name": info.get("site_name", "VideoGen.AI"),
|
|
"site_logo": to_full_url(info.get("site_logo")),
|
|
"user_agreement_privacy_url": to_full_url(info.get("user_agreement_privacy_url")),
|
|
"site_copyright": info.get("site_copyright", "© 2024 民众智创 版权所有"),
|
|
}
|
|
|
|
|
|
@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)}
|