235 lines
7.8 KiB
Python
235 lines
7.8 KiB
Python
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||
from fastapi.responses import Response
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
from sqlalchemy import select, and_
|
||
from fastcrud import FastCRUD
|
||
from datetime import timedelta, datetime
|
||
import redis.asyncio as redis
|
||
|
||
from database import get_db
|
||
from models import User
|
||
from schemas import (
|
||
UserCreate,
|
||
UserResponse,
|
||
UserLogin,
|
||
Token,
|
||
UserCreateRequest,
|
||
UserUpdate,
|
||
)
|
||
from utils.password import verify_password, get_password_hash
|
||
from utils.jwt import create_access_token
|
||
from config import app_settings
|
||
from utils.redis_client import get_redis, set_token_in_redis
|
||
from utils.captcha import create_captcha, validate_captcha
|
||
from middleware import get_current_user
|
||
from utils.like_filter import ilike_contains
|
||
|
||
router = APIRouter(prefix="/users", tags=["users"])
|
||
|
||
# 创建FastCRUD实例
|
||
user_crud = FastCRUD(User, UserResponse)
|
||
|
||
|
||
@router.get("/captcha")
|
||
async def get_captcha(redis_conn: redis.Redis = Depends(get_redis)):
|
||
"""获取验证码图片"""
|
||
# 创建验证码
|
||
captcha_id, image_bytes = await create_captcha(redis_conn)
|
||
|
||
# 返回验证码图片,同时在响应头中返回验证码ID
|
||
response = Response(content=image_bytes, media_type="image/png")
|
||
response.headers["X-Captcha-ID"] = captcha_id
|
||
|
||
return response
|
||
|
||
|
||
@router.post("/login", response_model=Token)
|
||
async def login(
|
||
user_data: UserLogin,
|
||
db: AsyncSession = Depends(get_db),
|
||
redis_conn: redis.Redis = Depends(get_redis),
|
||
):
|
||
"""用户登录"""
|
||
# 验证验证码
|
||
is_captcha_valid = await validate_captcha(
|
||
redis_conn, user_data.captcha_id, user_data.captcha_code
|
||
)
|
||
if not is_captcha_valid:
|
||
raise HTTPException(
|
||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||
detail="验证码错误或已过期",
|
||
headers={"WWW-Authenticate": "Bearer"},
|
||
)
|
||
|
||
# 查找用户
|
||
result = await db.execute(select(User).filter(User.username == user_data.username))
|
||
user = result.scalar_one_or_none()
|
||
|
||
# 验证用户是否存在且密码正确
|
||
if not user or not verify_password(user_data.password, user.password_hash):
|
||
raise HTTPException(
|
||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||
detail="用户名或密码错误",
|
||
headers={"WWW-Authenticate": "Bearer"},
|
||
)
|
||
|
||
# 创建访问令牌
|
||
access_token_expires = timedelta(minutes=app_settings.access_token_expire_minutes)
|
||
access_token = create_access_token(
|
||
data={"sub": str(user.id), "username": user.username},
|
||
expires_delta=access_token_expires,
|
||
)
|
||
|
||
# 将令牌存储到Redis
|
||
expire_seconds = int(access_token_expires.total_seconds())
|
||
await set_token_in_redis(redis_conn, user.id, access_token, expire_seconds)
|
||
|
||
# 更新用户最后登录时间
|
||
user.lastlogintime = datetime.utcnow()
|
||
await db.commit()
|
||
|
||
return {"access_token": access_token, "token_type": "bearer"}
|
||
|
||
|
||
@router.post("/", response_model=UserResponse)
|
||
async def create_user(user_data: UserCreateRequest, db: AsyncSession = Depends(get_db)):
|
||
"""创建用户"""
|
||
# 检查用户名是否已存在
|
||
result = await db.execute(select(User).filter(User.username == user_data.username))
|
||
existing_user = result.scalar_one_or_none()
|
||
if existing_user:
|
||
raise HTTPException(status_code=400, detail="用户名已存在")
|
||
|
||
# 生成密码哈希值
|
||
hashed_password = get_password_hash(user_data.password)
|
||
|
||
# 准备创建数据
|
||
create_data = user_data.model_dump()
|
||
create_data["password_hash"] = hashed_password
|
||
del create_data["password"] # 删除明文密码
|
||
# 转换为UserCreate模型
|
||
user_data_create = UserCreate(**create_data)
|
||
|
||
# 使用FastCRUD创建用户
|
||
created_user = await user_crud.create(
|
||
db, user_data_create, schema_to_select=UserResponse, return_as_model=True
|
||
)
|
||
return created_user
|
||
|
||
|
||
@router.get("/", response_model=list[UserResponse])
|
||
async def get_users(
|
||
db: AsyncSession = Depends(get_db),
|
||
current_user: User = Depends(get_current_user),
|
||
username: str | None = Query(None, description="登录名,模糊匹配"),
|
||
nickname: str | None = Query(None, description="昵称,模糊匹配"),
|
||
email: str | None = Query(None, description="邮箱,模糊匹配"),
|
||
phone: str | None = Query(None, description="手机,模糊匹配"),
|
||
):
|
||
"""分页列表用户;支持登录名、昵称、邮箱、手机的模糊查询与组合查询(AND)"""
|
||
stmt = select(User)
|
||
filters = []
|
||
for col, raw in (
|
||
(User.username, username),
|
||
(User.nickname, nickname),
|
||
(User.email, email),
|
||
(User.phone, phone),
|
||
):
|
||
cond = ilike_contains(col, raw)
|
||
if cond is not None:
|
||
filters.append(cond)
|
||
if filters:
|
||
stmt = stmt.where(and_(*filters))
|
||
result = await db.execute(stmt)
|
||
rows = result.scalars().all()
|
||
return [UserResponse.model_validate(u) for u in rows]
|
||
|
||
|
||
@router.get("/me", response_model=UserResponse)
|
||
async def get_current_user_info(current_user: User = Depends(get_current_user)):
|
||
"""获取当前用户信息(须声明在 /{user_id} 之前,否则会被当成 id)"""
|
||
return current_user
|
||
|
||
|
||
@router.get("/{user_id}", response_model=UserResponse)
|
||
async def get_user(
|
||
user_id: int,
|
||
db: AsyncSession = Depends(get_db),
|
||
current_user: User = Depends(get_current_user),
|
||
):
|
||
"""获取单个用户"""
|
||
user = await user_crud.get(db, id=user_id)
|
||
if not user:
|
||
raise HTTPException(status_code=404, detail="用户不存在")
|
||
return user
|
||
|
||
|
||
@router.put("/{user_id}", response_model=UserResponse)
|
||
async def update_user(
|
||
user_id: int,
|
||
user_data: UserUpdate,
|
||
db: AsyncSession = Depends(get_db),
|
||
current_user: User = Depends(get_current_user),
|
||
):
|
||
"""更新用户"""
|
||
# 检查用户名是否被其他用户使用
|
||
result = await db.execute(
|
||
select(User).filter(User.username == user_data.username, User.id != user_id)
|
||
)
|
||
existing_user = result.scalar_one_or_none()
|
||
if existing_user:
|
||
raise HTTPException(status_code=400, detail="用户名已存在")
|
||
|
||
# 生成密码哈希值
|
||
hashed_password = (
|
||
get_password_hash(user_data.password) if user_data.password else None
|
||
)
|
||
|
||
# 准备更新数据
|
||
update_data = user_data.model_dump()
|
||
if hashed_password:
|
||
update_data["password_hash"] = hashed_password
|
||
del update_data["password"] # 删除明文密码
|
||
|
||
# 使用FastCRUD更新用户
|
||
updated_user = await user_crud.update(
|
||
db, update_data, id=user_id, schema_to_select=UserResponse, return_as_model=True
|
||
)
|
||
if not updated_user:
|
||
raise HTTPException(status_code=404, detail="用户不存在")
|
||
return updated_user
|
||
|
||
|
||
@router.delete("/{user_id}")
|
||
async def delete_user(
|
||
user_id: int,
|
||
db: AsyncSession = Depends(get_db),
|
||
current_user: User = Depends(get_current_user),
|
||
):
|
||
"""删除用户"""
|
||
# 检查用户是否存在
|
||
result = await db.execute(select(User).filter(User.id == user_id))
|
||
existing_user = result.scalar_one_or_none()
|
||
if not existing_user:
|
||
raise HTTPException(status_code=404, detail="用户不存在")
|
||
elif existing_user.id == current_user.id:
|
||
raise HTTPException(status_code=400, detail="不能删除当前登录用户")
|
||
elif existing_user.role_id == 1:
|
||
raise HTTPException(status_code=400, detail="不能删除管理员用户")
|
||
else:
|
||
# 使用FastCRUD删除用户
|
||
deleted = await user_crud.delete(db, id=user_id)
|
||
return {"message": "用户删除成功"}
|
||
|
||
|
||
@router.post("/logout")
|
||
async def logout(
|
||
current_user: User = Depends(get_current_user),
|
||
redis_conn: redis.Redis = Depends(get_redis),
|
||
):
|
||
"""用户登出"""
|
||
from utils.redis_client import delete_token_from_redis
|
||
|
||
await delete_token_from_redis(redis_conn, current_user.id)
|
||
return {"message": "登出成功"}
|