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) # 计算 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, 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() 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", "operation_manual" ])) ) 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 民众智创 版权所有"), "operation_manual": info.get("operation_manual", ""), } @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)}