2be0778ec7
- 后端:auth 路由(register/login/refresh/me)+ JWT + bcrypt 密码哈希 - 依赖注入:get_current_user + require_role 角色权限校验 - 跨数据库兼容:JSONBType(PG 用 JSONB,SQLite 用 JSON) - 测试:11 个认证测试 + 4 个健康检查测试 = 15 passed
164 lines
4.8 KiB
Python
164 lines
4.8 KiB
Python
"""认证路由:登录 / 注册 / 刷新 token / 获取当前用户。"""
|
|
|
|
from fastapi import APIRouter, Depends, HTTPException, status
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.core.config import settings
|
|
from app.core.database import get_db
|
|
from app.core.dependencies import get_current_user
|
|
from app.core.security import (
|
|
create_access_token,
|
|
create_refresh_token,
|
|
decode_token,
|
|
hash_password,
|
|
verify_password,
|
|
)
|
|
from app.models.tenant import Tenant
|
|
from app.models.user import User
|
|
from app.schemas.auth import (
|
|
LoginRequest,
|
|
RefreshRequest,
|
|
RegisterRequest,
|
|
TokenResponse,
|
|
UserInfo,
|
|
)
|
|
from app.schemas.common import ApiResponse, success
|
|
|
|
router = APIRouter(prefix="/auth", tags=["auth"])
|
|
|
|
|
|
@router.post("/register", response_model=ApiResponse[TokenResponse])
|
|
async def register(req: RegisterRequest, db: AsyncSession = Depends(get_db)):
|
|
"""用户注册(创建新租户 + 首个用户)。"""
|
|
# 检查邮箱是否已存在
|
|
existing = await db.execute(select(User).where(User.email == req.email))
|
|
if existing.scalar_one_or_none():
|
|
raise HTTPException(
|
|
status_code=status.HTTP_409_CONFLICT,
|
|
detail="邮箱已注册",
|
|
)
|
|
|
|
# 创建租户
|
|
tenant = Tenant(name=req.tenant_name, type="vc")
|
|
db.add(tenant)
|
|
await db.flush()
|
|
|
|
# 创建用户
|
|
user = User(
|
|
tenant_id=tenant.id,
|
|
email=req.email,
|
|
password_hash=hash_password(req.password),
|
|
name=req.name,
|
|
role=req.role,
|
|
)
|
|
db.add(user)
|
|
await db.flush()
|
|
|
|
access_token = create_access_token(
|
|
subject=user.id,
|
|
extra_claims={"role": user.role, "tenant_id": user.tenant_id},
|
|
)
|
|
refresh_token = create_refresh_token(subject=user.id)
|
|
|
|
return success(
|
|
data=TokenResponse(
|
|
access_token=access_token,
|
|
refresh_token=refresh_token,
|
|
expires_in=settings.jwt_access_token_ttl_minutes * 60,
|
|
),
|
|
message="注册成功",
|
|
)
|
|
|
|
|
|
@router.post("/login", response_model=ApiResponse[TokenResponse])
|
|
async def login(req: LoginRequest, db: AsyncSession = Depends(get_db)):
|
|
"""用户登录。"""
|
|
result = await db.execute(select(User).where(User.email == req.email))
|
|
user = result.scalar_one_or_none()
|
|
|
|
if not user or not verify_password(req.password, user.password_hash):
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="邮箱或密码错误",
|
|
)
|
|
|
|
if not user.is_active:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="账户已禁用",
|
|
)
|
|
|
|
access_token = create_access_token(
|
|
subject=user.id,
|
|
extra_claims={"role": user.role, "tenant_id": user.tenant_id},
|
|
)
|
|
refresh_token = create_refresh_token(subject=user.id)
|
|
|
|
return success(
|
|
data=TokenResponse(
|
|
access_token=access_token,
|
|
refresh_token=refresh_token,
|
|
expires_in=settings.jwt_access_token_ttl_minutes * 60,
|
|
),
|
|
message="登录成功",
|
|
)
|
|
|
|
|
|
@router.post("/refresh", response_model=ApiResponse[TokenResponse])
|
|
async def refresh_token(req: RefreshRequest, db: AsyncSession = Depends(get_db)):
|
|
"""刷新 access token。"""
|
|
try:
|
|
payload = decode_token(req.refresh_token)
|
|
except Exception:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="无效的 refresh token",
|
|
)
|
|
|
|
if payload.get("type") != "refresh":
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="无效的 token 类型",
|
|
)
|
|
|
|
user_id = payload.get("sub")
|
|
result = await db.execute(select(User).where(User.id == user_id))
|
|
user = result.scalar_one_or_none()
|
|
|
|
if not user or not user.is_active:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="用户不存在或已禁用",
|
|
)
|
|
|
|
access_token = create_access_token(
|
|
subject=user.id,
|
|
extra_claims={"role": user.role, "tenant_id": user.tenant_id},
|
|
)
|
|
new_refresh_token = create_refresh_token(subject=user.id)
|
|
|
|
return success(
|
|
data=TokenResponse(
|
|
access_token=access_token,
|
|
refresh_token=new_refresh_token,
|
|
expires_in=settings.jwt_access_token_ttl_minutes * 60,
|
|
),
|
|
message="刷新成功",
|
|
)
|
|
|
|
|
|
@router.get("/me", response_model=ApiResponse[UserInfo])
|
|
async def get_me(user: User = Depends(get_current_user)):
|
|
"""获取当前用户信息。"""
|
|
return success(
|
|
data=UserInfo(
|
|
id=user.id,
|
|
email=user.email,
|
|
name=user.name,
|
|
role=user.role,
|
|
tenant_id=user.tenant_id,
|
|
is_active=user.is_active,
|
|
),
|
|
)
|