Files
video-gen/video-gen-api/app/api/v1/auth.py
T

442 lines
16 KiB
Python

from datetime import datetime, timezone, timedelta
CST = timezone(timedelta(hours=8))
from fastapi import APIRouter, Depends, HTTPException, Request, 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,
security,
)
from fastapi.security import HTTPAuthorizationCredentials
from app.models.system_config import SystemConfig
from app.models.user import User
from app.schemas.auth import (
ChangePasswordRequest,
ChangeUsernameRequest,
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.utils.device import detect_device_type
from app.services.sms import verify_sms_code
from app.services.resource_capacity_service import get_user_resource_capacity_usage
from app.enums.credit_balance import CreditBalanceSourceType, CreditLevel
from app.services.credit.ledger_service import grant_credits
from app.services.credit.query_service import attach_credit_snapshot, get_available_credits
from app.services.credit.time_policy import add_natural_months
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, device_type: str = "pc") -> dict:
user.credits = round(user.credits, 2)
# 按设备类型递增对应版本号 — 单设备登录(同端互斥)
if device_type == "mobile":
user.mobile_token_version = (user.mobile_token_version or 0) + 1
version = user.mobile_token_version
else:
user.pc_token_version = (user.pc_token_version or 0) + 1
version = user.pc_token_version
token = create_access_token(user.id, remember_me, version, device_type)
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:
attach_credit_snapshot(user, 0)
return
now = datetime.now(timezone.utc)
result = await grant_credits(
db,
user_id=user.id,
amount=credits,
description=f"注册赠送 {credits} 积分",
source_type=CreditBalanceSourceType.REGISTER_GIFT.value,
source_id=user.id,
valid_from=now,
expires_at=add_natural_months(now, 1),
credit_level=CreditLevel.PROMOTIONAL.value,
related_id=user.id,
biz_key=f"register-gift:{user.id}",
request_time=now,
)
attach_credit_snapshot(user, result.balance_after)
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)
)
if enabled_result.scalar_one_or_none() != "true":
attach_credit_snapshot(user, await get_available_credits(db, user.id))
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")
now_cst = datetime.now(CST)
if credits > 0:
next_midnight_cst = datetime.combine(now_cst.date() + timedelta(days=1), datetime.min.time(), tzinfo=CST)
result = await grant_credits(
db,
user_id=user.id,
amount=credits,
description=f"每日登录赠送 {credits} 积分",
source_type=CreditBalanceSourceType.DAILY_LOGIN.value,
source_id=now_cst.date().isoformat(),
valid_from=now_cst,
expires_at=next_midnight_cst,
credit_level=CreditLevel.PROMOTIONAL.value,
related_id=user.id,
biz_key=f"daily-login:{user.id}:{now_cst.date().isoformat()}",
request_time=now_cst,
)
attach_credit_snapshot(user, result.balance_after)
else:
attach_credit_snapshot(user, await get_available_credits(db, user.id))
@router.post(
"/login",
summary="客户端密码登录",
description="保留原有用户名/手机号 + 密码登录。仅允许 frontend 用户登录;管理员仍使用 /auth/admin-login。",
)
async def login(req: LoginRequest, request: Request, 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(CST)
await db.flush()
device_type = detect_device_type(request.headers.get("user-agent"))
return _token_response(user, req.remember_me, device_type)
@router.post(
"/sms-login",
summary="客户端短信验证码登录",
description="新增兼容登录方式:手机号 + 短信验证码登录。不覆盖 /auth/login 密码登录。仅允许 frontend 用户登录。",
)
async def sms_login(req: SmsLoginRequest, request: Request, 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(CST)
await db.flush()
device_type = detect_device_type(request.headers.get("user-agent"))
return _token_response(user, req.remember_me, device_type)
@router.post(
"/register",
summary="客户端手机号短信注册",
description="手机号 + 注册短信验证码注册。注册成功后 username 默认等于手机号,不生成密码;前端需根据 must_set_password 引导用户设置密码。",
)
async def register(req: RegisterRequest, request: Request, 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(CST),
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)
# 按设备类型递增对应版本号
device_type = detect_device_type(request.headers.get("user-agent"))
if device_type == "mobile":
user.mobile_token_version = (user.mobile_token_version or 0) + 1
version = user.mobile_token_version
else:
user.pc_token_version = (user.pc_token_version or 0) + 1
version = user.pc_token_version
token = create_access_token(user.id, False, version, device_type)
return {
"access_token": token,
"token_type": "bearer",
"user": UserOut.model_validate(user),
"must_set_password": True,
}
@router.post("/logout")
async def logout(
credentials: HTTPAuthorizationCredentials | None = Depends(security),
current_user: User = Depends(get_current_user_allow_password_pending),
db: AsyncSession = Depends(get_db),
):
# 按设备类型递增对应版本号 — 使当前 token 失效
device_type = "pc"
if credentials:
payload = decode_access_token(credentials.credentials)
if payload:
device_type = payload.get("dev", "pc")
if device_type == "mobile":
current_user.mobile_token_version = (current_user.mobile_token_version or 0) + 1
else:
current_user.pc_token_version = (current_user.pc_token_version or 0) + 1
await db.flush()
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)
# 计算 is_team_manager 和 team_name
is_team_manager = False
team_name = None
if current_user.team_id:
from app.services.team_manager_service import is_team_manager
from app.services.team_service import batch_get_team_name_map
is_team_manager = await is_team_manager(db, current_user.id, current_user.team_id)
name_map = await batch_get_team_name_map(db, [current_user.team_id])
team_name = name_map.get(current_user.team_id)
return UserOut.model_validate(current_user).model_copy(
update={
"resource_capacity": resource_capacity,
"is_team_manager": is_team_manager,
"team_id": current_user.team_id,
"team_name": team_name,
}
)
@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(CST)
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(CST)
await db.flush()
return {"message": "密码修改成功"}
@router.post("/change-username")
async def change_username(
req: ChangeUsernameRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
current_user.username = req.new_username.strip()
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", "operation_manual", "login_bg_video", "optimize_hold_credits", "site_banner", "site_banner_version"
]))
)
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", "智创"),
"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", "© 2026 智创 版权所有"),
"operation_manual": info.get("operation_manual", ""),
"login_bg_video": to_full_url(info.get("login_bg_video")) if info.get("login_bg_video") else "",
"optimize_hold_credits": int(info.get("optimize_hold_credits") or 5),
"site_banner": info.get("site_banner", ""),
"site_banner_version": int(info.get("site_banner_version") or 0),
}
@router.post("/admin-login")
async def admin_login(req: LoginRequest, request: Request, 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(CST)
await db.flush()
device_type = detect_device_type(request.headers.get("user-agent"))
token = create_access_token(user.id, req.remember_me, 0, device_type)
user.credits = round(user.credits, 2)
return {"access_token": token, "token_type": "bearer", "user": UserOut.model_validate(user)}