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)}