feat(auth): 重构用户认证模块并优化初始化数据逻辑

- 重构用户和角色模型,优化字段定义和关系
- 增强初始化数据脚本,支持数据更新检查
- 改进用户和角色API端点,增加验证逻辑
- 扩展Pydantic模型,分离请求和响应模式
- 自定义Swagger UI界面并优化API文档
- 移除测试文件并更新依赖项配置
This commit is contained in:
jayhgq 2026-03-01 23:54:41 +08:00
parent f1da855793
commit 001d0abca2
27 changed files with 467 additions and 176 deletions

View File

@ -14,52 +14,97 @@ role_crud = FastCRUD(Role, RoleResponse)
@router.post("/", response_model=RoleResponse) @router.post("/", response_model=RoleResponse)
async def create_role(role_data: RoleCreate, db: AsyncSession = Depends(get_db), current_user: User = Depends(get_current_user)): async def create_role(
role_data: RoleCreate,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user),
):
"""创建角色""" """创建角色"""
# 创建角色实例 # 检查角色名是否已存在
db_role = Role( from sqlalchemy import select
name=role_data.name,
creator=role_data.creator result = await db.execute(select(Role).filter(Role.name == role_data.name))
) existing_role = result.scalar_one_or_none()
if existing_role:
raise HTTPException(status_code=400, detail="角色名已存在")
# 准备创建数据
create_data = role_data.model_dump()
# 使用FastCRUD创建角色 # 使用FastCRUD创建角色
created_role = await role_crud.create(db, db_role) created_role = await role_crud.create(db, create_data)
return created_role return created_role
@router.get("/", response_model=list[RoleResponse]) @router.get("/", response_model=list[RoleResponse])
async def get_roles(db: AsyncSession = Depends(get_db), current_user: User = Depends(get_current_user)): async def get_roles(
db: AsyncSession = Depends(get_db), current_user: User = Depends(get_current_user)
):
"""获取所有角色""" """获取所有角色"""
roles = await role_crud.get_multi(db) result = await role_crud.get_multi(db)
return roles return result["data"]
@router.get("/{role_id}", response_model=RoleResponse) @router.get("/{role_id}", response_model=RoleResponse)
async def get_role(role_id: int, db: AsyncSession = Depends(get_db), current_user: User = Depends(get_current_user)): async def get_role(
role_id: int,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user),
):
"""获取单个角色""" """获取单个角色"""
role = await role_crud.get(db, role_id) role = await role_crud.get(db, {"id": role_id})
if not role: if not role:
raise HTTPException(status_code=404, detail="角色不存在") raise HTTPException(status_code=404, detail="角色不存在")
return role return role
@router.put("/{role_id}", response_model=RoleResponse) @router.put("/{role_id}", response_model=RoleResponse)
async def update_role(role_id: int, role_data: RoleCreate, db: AsyncSession = Depends(get_db), current_user: User = Depends(get_current_user)): async def update_role(
role_id: int,
role_data: RoleCreate,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user),
):
"""更新角色""" """更新角色"""
# 检查角色名是否被其他角色使用
from sqlalchemy import select
result = await db.execute(
select(Role).filter(Role.name == role_data.name, Role.id != role_id)
)
existing_role = result.scalar_one_or_none()
if existing_role:
raise HTTPException(status_code=400, detail="角色名已存在")
# 准备更新数据 # 准备更新数据
update_data = role_data.dict() update_data = role_data.model_dump()
# 使用FastCRUD更新角色 # 使用FastCRUD更新角色
updated_role = await role_crud.update(db, role_id, update_data) updated_role = await role_crud.update(db, {"id": role_id}, update_data)
if not updated_role: if not updated_role:
raise HTTPException(status_code=404, detail="角色不存在") raise HTTPException(status_code=404, detail="角色不存在")
return updated_role return updated_role
@router.delete("/{role_id}") @router.delete("/{role_id}")
async def delete_role(role_id: int, db: AsyncSession = Depends(get_db), current_user: User = Depends(get_current_user)): async def delete_role(
role_id: int,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user),
):
"""删除角色""" """删除角色"""
deleted = await role_crud.delete(db, role_id) # 检查是否有用户使用此角色
from sqlalchemy import select
result = await db.execute(select(User).filter(User.role_id == role_id))
users = result.scalars().all()
if users:
raise HTTPException(
status_code=400, detail=f"该角色正在被 {len(users)} 个用户使用,无法删除"
)
# 使用FastCRUD删除角色
deleted = await role_crud.delete(db, {"id": role_id})
if not deleted: if not deleted:
raise HTTPException(status_code=404, detail="角色不存在") raise HTTPException(status_code=404, detail="角色不存在")
return {"message": "角色删除成功"} return {"message": "角色删除成功"}

View File

@ -1,12 +1,20 @@
from fastapi import APIRouter, Depends, HTTPException, status from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from fastcrud import FastCRUD from fastcrud import FastCRUD
from datetime import timedelta from datetime import timedelta, datetime
import redis.asyncio as redis import redis.asyncio as redis
from database import get_db from database import get_db
from models import User from models import User
from schemas import UserCreate, UserResponse, UserLogin, Token from schemas import (
UserCreate,
UserResponse,
UserLogin,
Token,
UserCreateRequest,
UserUpdate,
)
from utils.password import verify_password, get_password_hash from utils.password import verify_password, get_password_hash
from utils.jwt import create_access_token from utils.jwt import create_access_token
from config import app_settings from config import app_settings
@ -20,11 +28,16 @@ user_crud = FastCRUD(User, UserResponse)
@router.post("/login", response_model=Token) @router.post("/login", response_model=Token)
async def login(user_data: UserLogin, db: AsyncSession = Depends(get_db), redis_conn: redis.Redis = Depends(get_redis)): async def login(
user_data: UserLogin,
db: AsyncSession = Depends(get_db),
redis_conn: redis.Redis = Depends(get_redis),
):
"""用户登录""" """用户登录"""
# 查找用户 # 查找用户
user = await db.query(User).filter(User.username == user_data.username).first() 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): if not user or not verify_password(user_data.password, user.password_hash):
raise HTTPException( raise HTTPException(
@ -32,91 +45,139 @@ async def login(user_data: UserLogin, db: AsyncSession = Depends(get_db), redis_
detail="用户名或密码错误", detail="用户名或密码错误",
headers={"WWW-Authenticate": "Bearer"}, headers={"WWW-Authenticate": "Bearer"},
) )
# 创建访问令牌 # 创建访问令牌
access_token_expires = timedelta(minutes=app_settings.access_token_expire_minutes) access_token_expires = timedelta(minutes=app_settings.access_token_expire_minutes)
access_token = create_access_token( access_token = create_access_token(
data={"sub": str(user.id), "username": user.username}, data={"sub": str(user.id), "username": user.username},
expires_delta=access_token_expires expires_delta=access_token_expires,
) )
# 将令牌存储到Redis # 将令牌存储到Redis
expire_seconds = int(access_token_expires.total_seconds()) expire_seconds = int(access_token_expires.total_seconds())
await set_token_in_redis(redis_conn, user.id, access_token, expire_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"} return {"access_token": access_token, "token_type": "bearer"}
@router.post("/", response_model=UserResponse) @router.post("/", response_model=UserResponse)
async def create_user(user_data: UserCreate, db: AsyncSession = Depends(get_db)): 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) hashed_password = get_password_hash(user_data.password)
# 创建用户实例,使用密码哈希值 # 准备创建数据
db_user = User( create_data = user_data.model_dump()
username=user_data.username, create_data["password_hash"] = hashed_password
nickname=user_data.nickname, del create_data["password"] # 删除明文密码
email=user_data.email, # 转换为UserCreate模型
phone=user_data.phone, user_data_create = UserCreate(**create_data)
wx_openid=user_data.wx_openid,
avatar=user_data.avatar,
password_hash=hashed_password,
role_id=user_data.role_id
)
# 使用FastCRUD创建用户 # 使用FastCRUD创建用户
created_user = await user_crud.create(db, db_user) created_user = await user_crud.create(
db, user_data_create, schema_to_select=UserResponse, return_as_model=True
)
return created_user return created_user
@router.get("/", response_model=list[UserResponse]) @router.get("/", response_model=list[UserResponse])
async def get_users(db: AsyncSession = Depends(get_db), current_user: User = Depends(get_current_user)): async def get_users(
db: AsyncSession = Depends(get_db), current_user: User = Depends(get_current_user)
):
"""获取所有用户""" """获取所有用户"""
users = await user_crud.get_multi(db) result = await user_crud.get_multi(db)
return users return result["data"]
@router.get("/{user_id}", response_model=UserResponse) @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)): 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, user_id) user = await user_crud.get(db, id=user_id)
if not user: if not user:
raise HTTPException(status_code=404, detail="用户不存在") raise HTTPException(status_code=404, detail="用户不存在")
return user return user
@router.put("/{user_id}", response_model=UserResponse) @router.put("/{user_id}", response_model=UserResponse)
async def update_user(user_id: int, user_data: UserCreate, db: AsyncSession = Depends(get_db), current_user: User = Depends(get_current_user)): 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) hashed_password = (
get_password_hash(user_data.password) if user_data.password else None
)
# 准备更新数据 # 准备更新数据
update_data = user_data.dict() update_data = user_data.model_dump()
update_data["password_hash"] = hashed_password if hashed_password:
update_data["password_hash"] = hashed_password
del update_data["password"] # 删除明文密码 del update_data["password"] # 删除明文密码
# 使用FastCRUD更新用户 # 使用FastCRUD更新用户
updated_user = await user_crud.update(db, user_id, update_data) updated_user = await user_crud.update(
db, update_data, id=user_id, schema_to_select=UserResponse, return_as_model=True
)
if not updated_user: if not updated_user:
raise HTTPException(status_code=404, detail="用户不存在") raise HTTPException(status_code=404, detail="用户不存在")
return updated_user return updated_user
@router.delete("/{user_id}") @router.delete("/{user_id}")
async def delete_user(user_id: int, db: AsyncSession = Depends(get_db), current_user: User = Depends(get_current_user)): async def delete_user(
user_id: int,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user),
):
"""删除用户""" """删除用户"""
deleted = await user_crud.delete(db, user_id) # 检查用户是否存在
if not deleted: 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="用户不存在") raise HTTPException(status_code=404, detail="用户不存在")
return {"message": "用户删除成功"} 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") @router.post("/logout")
async def logout(current_user: User = Depends(get_current_user), redis_conn: redis.Redis = Depends(get_redis)): 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 from utils.redis_client import delete_token_from_redis
await delete_token_from_redis(redis_conn, current_user.id) await delete_token_from_redis(redis_conn, current_user.id)
return {"message": "登出成功"} return {"message": "登出成功"}
@ -124,4 +185,5 @@ async def logout(current_user: User = Depends(get_current_user), redis_conn: red
@router.get("/me", response_model=UserResponse) @router.get("/me", response_model=UserResponse)
async def get_current_user_info(current_user: User = Depends(get_current_user)): async def get_current_user_info(current_user: User = Depends(get_current_user)):
"""获取当前用户信息""" """获取当前用户信息"""
return current_user print(current_user.id)
return user

View File

@ -5,6 +5,7 @@ from urllib.parse import quote_plus
class PostgreSQLSettings(BaseSettings): class PostgreSQLSettings(BaseSettings):
"""PostgreSQL数据库配置""" """PostgreSQL数据库配置"""
host: str = "127.0.0.1" host: str = "127.0.0.1"
port: int = 5432 port: int = 5432
user: str = "admin" user: str = "admin"
@ -12,44 +13,46 @@ class PostgreSQLSettings(BaseSettings):
database: str = "booksystem" database: str = "booksystem"
min_size: int = 5 min_size: int = 5
max_size: int = 20 max_size: int = 20
@property @property
def url(self) -> str: def url(self) -> str:
"""生成数据库连接URL""" """生成数据库连接URL"""
encoded_password = quote_plus(self.password) encoded_password = quote_plus(self.password)
return f"postgresql+asyncpg://{self.user}:{encoded_password}@{self.host}:{self.port}/{self.database}" return f"postgresql+asyncpg://{self.user}:{encoded_password}@{self.host}:{self.port}/{self.database}"
class Config: class Config:
env_prefix = "POSTGRES_" env_prefix = "POSTGRES_"
class RedisSettings(BaseSettings): class RedisSettings(BaseSettings):
"""Redis配置""" """Redis配置"""
host: str = "127.0.0.1" host: str = "127.0.0.1"
port: int = 6379 port: int = 6379
password: Optional[str] = "RedisAdmin@123" password: Optional[str] = "RedisAdmin@123"
db: int = 0 db: int = 0
@property @property
def url(self) -> str: def url(self) -> str:
"""生成Redis连接URL""" """生成Redis连接URL"""
if self.password: if self.password:
return f"redis://:{self.password}@{self.host}:{self.port}/{self.db}" return f"redis://:{self.password}@{self.host}:{self.port}/{self.db}"
return f"redis://{self.host}:{self.port}/{self.db}" return f"redis://{self.host}:{self.port}/{self.db}"
class Config: class Config:
env_prefix = "REDIS_" env_prefix = "REDIS_"
class AppSettings(BaseSettings): class AppSettings(BaseSettings):
"""应用配置""" """应用配置"""
name: str = "Book System API" name: str = "Book System API"
version: str = "1.0.0" version: str = "1.0.0"
debug: bool = False debug: bool = False
secret_key: str = "BookSystemAPIQRAdmin123" secret_key: str = "BookSystemAPIQRAdmin123"
algorithm: str = "HS256" algorithm: str = "HS256"
access_token_expire_minutes: int = 30 access_token_expire_minutes: int = 10080
class Config: class Config:
env_prefix = "APP_" env_prefix = "APP_"
@ -57,4 +60,4 @@ class AppSettings(BaseSettings):
# 创建设置实例 # 创建设置实例
postgres_settings = PostgreSQLSettings() postgres_settings = PostgreSQLSettings()
redis_settings = RedisSettings() redis_settings = RedisSettings()
app_settings = AppSettings() app_settings = AppSettings()

View File

@ -6,36 +6,57 @@ from utils.password import get_password_hash
async def init_roles(db: AsyncSession): async def init_roles(db: AsyncSession):
"""初始化角色数据""" """初始化角色数据"""
# 检查是否已经存在角色数据 # 定义需要初始化的角色
result = await db.execute(select(Role)) required_roles = [
existing_roles = result.scalars().all() {"name": "管理员", "creator": "system"},
if existing_roles: {"name": "普通用户", "creator": "system"},
print("角色数据已存在,跳过初始化") {"name": "访客", "creator": "system"},
return {"name": "会员", "creator": "system"},
# 创建初始角色
roles = [
Role(name="管理员", creator="system"),
Role(name="普通用户", creator="system"),
Role(name="访客", creator="system"),
] ]
for role in roles: # 检查现有角色
db.add(role) result = await db.execute(select(Role))
existing_roles = result.scalars().all()
existing_role_names = {role.name for role in existing_roles}
await db.commit() # 检查需要创建或更新的角色
print("角色数据初始化完成") roles_to_add = []
roles_to_update = []
for role_data in required_roles:
role_name = role_data["name"]
if role_name not in existing_role_names:
# 角色不存在,需要创建
roles_to_add.append(Role(**role_data))
else:
# 角色存在,检查是否需要更新
existing_role = next(
role for role in existing_roles if role.name == role_name
)
if existing_role.creator != role_data["creator"]:
# 需要更新
existing_role.creator = role_data["creator"]
roles_to_update.append(existing_role)
# 执行创建和更新操作
if roles_to_add:
for role in roles_to_add:
db.add(role)
await db.commit()
print(f"创建了 {len(roles_to_add)} 个角色")
if roles_to_update:
await db.commit()
print(f"更新了 {len(roles_to_update)} 个角色")
if not roles_to_add and not roles_to_update:
print("角色数据已存在且无需更新,跳过初始化")
print("角色数据初始化/更新完成")
async def init_admin_user(db: AsyncSession): async def init_admin_user(db: AsyncSession):
"""初始化管理员用户""" """初始化管理员用户"""
# 检查是否已经存在管理员用户
result = await db.execute(select(User).filter(User.username == "admin"))
existing_admin = result.scalar_one_or_none()
if existing_admin:
print("管理员用户已存在,跳过初始化")
return
# 获取管理员角色 # 获取管理员角色
result = await db.execute(select(Role).filter(Role.name == "管理员")) result = await db.execute(select(Role).filter(Role.name == "管理员"))
admin_role = result.scalar_one_or_none() admin_role = result.scalar_one_or_none()
@ -43,32 +64,63 @@ async def init_admin_user(db: AsyncSession):
print("管理员角色不存在,请先初始化角色数据") print("管理员角色不存在,请先初始化角色数据")
return return
# 创建管理员用户 # 定义需要的管理员用户数据
admin_user = User( admin_data = {
username="admin", "username": "admin",
nickname="系统管理员", "nickname": "系统管理员",
email="admin@example.com", "email": "admin@example.com",
phone="13800138000", "phone": "13800138000",
password_hash=get_password_hash("admin"), "password_hash": get_password_hash("admin"),
role_id=admin_role.id, "role_id": admin_role.id,
avatar="", "avatar": "",
wx_openid="" "wx_openid": "",
) }
db.add(admin_user) # 检查管理员用户是否存在
await db.commit() result = await db.execute(select(User).filter(User.username == "admin"))
print("管理员用户初始化完成") existing_admin = result.scalar_one_or_none()
if not existing_admin:
# 管理员不存在,创建
admin_user = User(**admin_data)
db.add(admin_user)
await db.commit()
print("创建了管理员用户")
else:
# 管理员存在,检查是否需要更新
update_needed = False
# 检查字段是否需要更新(不包括密码,因为密码只在首次创建时设置)
if existing_admin.nickname != admin_data["nickname"]:
existing_admin.nickname = admin_data["nickname"]
update_needed = True
if existing_admin.email != admin_data["email"]:
existing_admin.email = admin_data["email"]
update_needed = True
if existing_admin.phone != admin_data["phone"]:
existing_admin.phone = admin_data["phone"]
update_needed = True
if existing_admin.role_id != admin_data["role_id"]:
existing_admin.role_id = admin_data["role_id"]
update_needed = True
if existing_admin.avatar != admin_data["avatar"]:
existing_admin.avatar = admin_data["avatar"]
update_needed = True
if existing_admin.wx_openid != admin_data["wx_openid"]:
existing_admin.wx_openid = admin_data["wx_openid"]
update_needed = True
if update_needed:
await db.commit()
print("更新了管理员用户")
else:
print("管理员用户已存在且无需更新,跳过初始化")
print("管理员用户初始化/更新完成")
async def init_test_user(db: AsyncSession): async def init_test_user(db: AsyncSession):
"""初始化测试用户""" """初始化测试用户"""
# 检查是否已经存在测试用户
result = await db.execute(select(User).filter(User.username == "test"))
existing_test_user = result.scalar_one_or_none()
if existing_test_user:
print("测试用户已存在,跳过初始化")
return
# 获取普通用户角色 # 获取普通用户角色
result = await db.execute(select(Role).filter(Role.name == "普通用户")) result = await db.execute(select(Role).filter(Role.name == "普通用户"))
user_role = result.scalar_one_or_none() user_role = result.scalar_one_or_none()
@ -76,21 +128,59 @@ async def init_test_user(db: AsyncSession):
print("普通用户角色不存在,请先初始化角色数据") print("普通用户角色不存在,请先初始化角色数据")
return return
# 创建测试用户 # 定义需要的测试用户数据
test_user = User( test_data = {
username="test", "username": "test",
nickname="测试用户", "nickname": "测试用户",
email="test@example.com", "email": "test@example.com",
phone="13800138001", "phone": "13800138001",
password_hash=get_password_hash("test"), "password_hash": get_password_hash("test"),
role_id=user_role.id, "role_id": user_role.id,
avatar="", "avatar": "",
wx_openid="" "wx_openid": "",
) }
db.add(test_user) # 检查测试用户是否存在
await db.commit() result = await db.execute(select(User).filter(User.username == "test"))
print("测试用户初始化完成") existing_test_user = result.scalar_one_or_none()
if not existing_test_user:
# 测试用户不存在,创建
test_user = User(**test_data)
db.add(test_user)
await db.commit()
print("创建了测试用户")
else:
# 测试用户存在,检查是否需要更新
update_needed = False
# 检查字段是否需要更新(不包括密码,因为密码只在首次创建时设置)
if existing_test_user.nickname != test_data["nickname"]:
existing_test_user.nickname = test_data["nickname"]
update_needed = True
if existing_test_user.email != test_data["email"]:
existing_test_user.email = test_data["email"]
update_needed = True
if existing_test_user.phone != test_data["phone"]:
existing_test_user.phone = test_data["phone"]
update_needed = True
if existing_test_user.role_id != test_data["role_id"]:
existing_test_user.role_id = test_data["role_id"]
update_needed = True
if existing_test_user.avatar != test_data["avatar"]:
existing_test_user.avatar = test_data["avatar"]
update_needed = True
if existing_test_user.wx_openid != test_data["wx_openid"]:
existing_test_user.wx_openid = test_data["wx_openid"]
update_needed = True
if update_needed:
await db.commit()
print("更新了测试用户")
else:
print("测试用户已存在且无需更新,跳过初始化")
print("测试用户初始化/更新完成")
async def init_all_data(db: AsyncSession): async def init_all_data(db: AsyncSession):

View File

@ -1,9 +1,12 @@
from fastapi import FastAPI from fastapi import FastAPI
from fastapi.staticfiles import StaticFiles
from fastapi.responses import HTMLResponse
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from apps.urls import api_router from apps.urls import api_router
from database import engine, get_db from database import engine, get_db
from models import Base from models import Base
from init_data import init_all_data from init_data import init_all_data
from pathlib import Path
@asynccontextmanager @asynccontextmanager
@ -12,39 +15,102 @@ async def lifespan(app: FastAPI):
async with engine.begin() as conn: async with engine.begin() as conn:
# 创建所有表 # 创建所有表
await conn.run_sync(Base.metadata.create_all) await conn.run_sync(Base.metadata.create_all)
# 初始化数据 # 初始化数据
async for db in get_db(): async for db in get_db():
await init_all_data(db) await init_all_data(db)
break break
yield yield
# 关闭时的清理工作 # 关闭时的清理工作
await engine.dispose() await engine.dispose()
# 实例化FastAPI # 实例化FastAPI
app = FastAPI( app = FastAPI(
title="Book System API", title="图书系统授权服务API",
description="图书系统API", description="""图书系统授权服务API是一套用于用户认证和授权的服务提供用户注册、登录、权限校验等功能。
同时实现了基于角色的访问控制RBAC支持自定义角色和权限还增加了日志记录功能方便监控和调试""",
version="1.0.0", version="1.0.0",
lifespan=lifespan lifespan=lifespan,
docs_url=None,
redoc_url="/redoc",
) )
# 挂载swagger-ui静态文件目录
swagger_ui_path = Path(__file__).parent.parent / "swagger-ui"
app.mount("/static", StaticFiles(directory=str(swagger_ui_path)), name="static")
# 自定义Swagger UI页面
@app.get("/docs", include_in_schema=False)
async def custom_swagger_ui_html():
html_content = """
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>Book System API - Swagger UI</title>
<link rel="stylesheet" type="text/css" href="/static/swagger-ui.css">
<link rel="icon" type="image/png" href="/static/favicon-32x32.png" sizes="32x32"/>
<style>
html {
box-sizing: border-box;
overflow: -moz-scrollbars-vertical;
overflow-y: scroll;
}
*, *:before, *:after {
box-sizing: inherit;
}
body {
margin: 0;
background: #fafafa;
}
</style>
</head>
<body>
<div id="swagger-ui"></div>
<script src="/static/swagger-ui-bundle.js"></script>
<script>
window.onload = function() {
const ui = SwaggerUIBundle({
url: "/openapi.json",
dom_id: '#swagger-ui',
deepLinking: true,
presets: [
SwaggerUIBundle.presets.apis,
SwaggerUIBundle.SwaggerUIStandalonePreset
],
layout: "BaseLayout",
persistAuthorization: true,
docExpansion: "list"
});
window.ui = ui;
};
</script>
</body>
</html>
"""
return HTMLResponse(content=html_content)
# 包含路由 # 包含路由
app.include_router(api_router) app.include_router(api_router)
# 声明装饰器方法和路径 # 声明装饰器方法和路径
@app.get("/") @app.get("/")
# 声明装饰器函数 # 声明装饰器函数
async def home(): async def home():
return {"message": "Hello World!!!"} return {"message": "Hello World!!!"}
# 如果使用命令行启动使用uvicorn 文件名:app --reload启动即可下面命令就不用写 # 如果使用命令行启动使用uvicorn 文件名:app --reload启动即可下面命令就不用写
# 如果写下面的命令就不需要命令行启动了直接用IDE运行即可 # 如果写下面的命令就不需要命令行启动了直接用IDE运行即可
if __name__ == "__main__": if __name__ == "__main__":
import uvicorn import uvicorn
import os import os
name = f"{os.path.splitext(os.path.basename(os.path.abspath(__file__)))[0]}:app" name = f"{os.path.splitext(os.path.basename(os.path.abspath(__file__)))[0]}:app"
uvicorn.run(name, host="0.0.0.0", port=8000, reload=True, reload_dirs=["_"]) uvicorn.run(name, host="0.0.0.0", port=8000, reload=True, reload_dirs=["_"])

View File

@ -1,6 +1,7 @@
from fastapi import Request, HTTPException, status, Depends from fastapi import Request, HTTPException, status, Depends
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
import redis.asyncio as redis import redis.asyncio as redis
from database import get_db from database import get_db
@ -10,10 +11,15 @@ from utils.redis_client import get_redis, check_token_in_redis
security = HTTPBearer() security = HTTPBearer()
async def get_current_user(request: Request, credentials: HTTPAuthorizationCredentials = Depends(security), db: AsyncSession = Depends(get_db), redis_conn: redis.Redis = Depends(get_redis)): async def get_current_user(
request: Request,
credentials: HTTPAuthorizationCredentials = Depends(security),
db: AsyncSession = Depends(get_db),
redis_conn: redis.Redis = Depends(get_redis),
):
"""获取当前用户""" """获取当前用户"""
token = credentials.credentials token = credentials.credentials
# 解码令牌 # 解码令牌
payload = decode_token(token) payload = decode_token(token)
if not payload: if not payload:
@ -22,9 +28,9 @@ async def get_current_user(request: Request, credentials: HTTPAuthorizationCrede
detail="无效的认证凭据", detail="无效的认证凭据",
headers={"WWW-Authenticate": "Bearer"}, headers={"WWW-Authenticate": "Bearer"},
) )
user_id = int(payload.get("sub")) user_id = int(payload.get("sub"))
# 检查令牌是否在Redis中 # 检查令牌是否在Redis中
is_valid = await check_token_in_redis(redis_conn, user_id, token) is_valid = await check_token_in_redis(redis_conn, user_id, token)
if not is_valid: if not is_valid:
@ -33,14 +39,13 @@ async def get_current_user(request: Request, credentials: HTTPAuthorizationCrede
detail="令牌已过期或已登出", detail="令牌已过期或已登出",
headers={"WWW-Authenticate": "Bearer"}, headers={"WWW-Authenticate": "Bearer"},
) )
# 查找用户 # 查找用户
from models import User from models import User
user = await db.query(User).filter(User.id == user_id).first()
result = await db.execute(select(User).filter(User.id == user_id))
user = result.scalar_one_or_none()
if not user: if not user:
raise HTTPException( raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="用户不存在")
status_code=status.HTTP_404_NOT_FOUND,
detail="用户不存在" return user
)
return user

View File

@ -8,19 +8,23 @@ Base = declarative_base()
class Role(Base): class Role(Base):
"""角色数据库模型""" """角色数据库模型"""
__tablename__ = "role" __tablename__ = "role"
id = Column(Integer, primary_key=True, index=True, comment="角色ID") id = Column(Integer, primary_key=True, index=True, comment="角色ID")
name = Column(String(50), nullable=False, unique=True, comment="角色名称") name = Column(String(50), nullable=False, unique=True, comment="角色名称")
createtime = Column(DateTime, default=datetime.utcnow, comment="创建时间") createtime = Column(DateTime, default=datetime.utcnow, comment="创建时间")
updatetime = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, comment="更新时间") updatetime = Column(
DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, comment="更新时间"
)
creator = Column(String(50), nullable=True, comment="创建人") creator = Column(String(50), nullable=True, comment="创建人")
# 关联到User # 关联到User
users = relationship("User", back_populates="role_rel") users = relationship("User", back_populates="role_rel")
class User(Base): class User(Base):
"""用户数据库模型""" """用户数据库模型"""
__tablename__ = "users" __tablename__ = "users"
id = Column(Integer, primary_key=True, index=True) id = Column(Integer, primary_key=True, index=True)
username = Column(String(50), unique=True, nullable=False, comment="登录名") username = Column(String(50), unique=True, nullable=False, comment="登录名")
@ -32,8 +36,10 @@ class User(Base):
avatar = Column(String(255), nullable=True, comment="头像路径") avatar = Column(String(255), nullable=True, comment="头像路径")
role_id = Column(Integer, ForeignKey("role.id"), default=2, comment="角色ID") role_id = Column(Integer, ForeignKey("role.id"), default=2, comment="角色ID")
createtime = Column(DateTime, default=datetime.utcnow, comment="创建时间") createtime = Column(DateTime, default=datetime.utcnow, comment="创建时间")
updatetime = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, comment="更新时间") updatetime = Column(
lastlogintime = Column(DateTime, nullable=True, comment="最后登录时间") DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, comment="更新时间"
)
lastlogintime = Column(DateTime, comment="最后登录时间")
# 关联到Role # 关联到Role
role_rel = relationship("Role", back_populates="users") role_rel = relationship("Role", back_populates="users")

View File

@ -6,6 +6,7 @@ fastapi_cdn_host==0.10.0
# Pydantic数据验证 # Pydantic数据验证
pydantic>=2.5.0 pydantic>=2.5.0
pydantic-settings>=2.1.0 pydantic-settings>=2.1.0
pydantic[email]>=2.5.0
# 密码加密 # 密码加密
passlib[bcrypt]>=1.7.4 passlib[bcrypt]>=1.7.4
@ -16,10 +17,12 @@ fastcrud==0.21.0
# JWT # JWT
python-jose[cryptography]>=3.3.0 python-jose[cryptography]>=3.3.0
# PostgreSQL + SQLAlchemy ORM包含异步支持 # SQLAlchemy ORM包含异步支持
sqlalchemy>=2.0.46 sqlalchemy>=2.0.46
asyncpg>=0.31.0
alembic>=1.18.4 alembic>=1.18.4
# PostgreSQL
asyncpg>=0.31.0
# Redis异步客户端 # Redis异步客户端
redis>=7.2.0 redis>=7.2.0

View File

@ -5,6 +5,7 @@ from datetime import datetime
class UserBase(BaseModel): class UserBase(BaseModel):
"""用户基础模型""" """用户基础模型"""
username: str = Field(..., min_length=3, max_length=50, description="登录名") username: str = Field(..., min_length=3, max_length=50, description="登录名")
nickname: str = Field(..., max_length=50, description="昵称") nickname: str = Field(..., max_length=50, description="昵称")
email: Optional[EmailStr] = Field(None, max_length=50, description="电子邮箱") email: Optional[EmailStr] = Field(None, max_length=50, description="电子邮箱")
@ -14,19 +15,43 @@ class UserBase(BaseModel):
role_id: int = Field(2, description="角色ID") role_id: int = Field(2, description="角色ID")
class UserCreateRequest(UserBase):
"""用户创建请求模型"""
password: str = Field(..., min_length=6, max_length=50, description="密码")
class UserCreate(UserBase): class UserCreate(UserBase):
"""用户创建模型""" """用户创建模型"""
password: str = Field(..., min_length=6, max_length=50, description="密码")
password_hash: str = Field(
..., min_length=6, max_length=128, description="密码哈希值"
)
class UserUpdate(BaseModel):
"""用户更新模型"""
username: str = Field(..., min_length=3, max_length=50, description="登录名")
nickname: str = Field(..., max_length=50, description="昵称")
email: Optional[EmailStr] = Field(None, max_length=50, description="电子邮箱")
phone: Optional[str] = Field(None, max_length=11, description="手机号码")
wx_openid: Optional[str] = Field(None, max_length=100, description="微信OpenID")
avatar: Optional[str] = Field(None, max_length=255, description="头像路径")
role_id: Optional[int] = Field(2, description="角色ID")
password: Optional[str] = Field(None, max_length=128, description="密码")
class UserLogin(BaseModel): class UserLogin(BaseModel):
"""用户登录模型""" """用户登录模型"""
username: str = Field(..., description="登录名") username: str = Field(..., description="登录名")
password: str = Field(..., description="密码") password: str = Field(..., description="密码")
class UserResponse(BaseModel): class UserResponse(BaseModel):
"""用户响应模型""" """用户响应模型"""
id: int = Field(..., description="用户ID") id: int = Field(..., description="用户ID")
username: str = Field(..., description="登录名") username: str = Field(..., description="登录名")
nickname: str = Field(..., description="昵称") nickname: str = Field(..., description="昵称")
@ -37,40 +62,45 @@ class UserResponse(BaseModel):
role_id: int = Field(..., description="角色ID") role_id: int = Field(..., description="角色ID")
createtime: datetime = Field(..., description="创建时间") createtime: datetime = Field(..., description="创建时间")
updatetime: datetime = Field(..., description="更新时间") updatetime: datetime = Field(..., description="更新时间")
lastlogintime: datetime = Field(..., description="最后登录时间") lastlogintime: Optional[datetime] = Field(None, description="最后登录时间")
class Config: class Config:
from_attributes = True from_attributes = True
class Token(BaseModel): class Token(BaseModel):
"""令牌模型""" """令牌模型"""
access_token: str = Field(..., description="访问令牌") access_token: str = Field(..., description="访问令牌")
token_type: str = Field(..., description="令牌类型") token_type: str = Field(..., description="令牌类型")
class TokenData(BaseModel): class TokenData(BaseModel):
"""令牌数据模型""" """令牌数据模型"""
user_id: Optional[int] = None user_id: Optional[int] = None
email: Optional[EmailStr] = None email: Optional[EmailStr] = None
class RoleBase(BaseModel): class RoleBase(BaseModel):
"""角色基础模型""" """角色基础模型"""
name: str = Field(..., max_length=50, description="角色名称") name: str = Field(..., max_length=50, description="角色名称")
class RoleCreate(RoleBase): class RoleCreate(RoleBase):
"""角色创建模型""" """角色创建模型"""
creator: Optional[str] = Field(None, max_length=50, description="创建人") creator: Optional[str] = Field(None, max_length=50, description="创建人")
class RoleResponse(RoleBase): class RoleResponse(RoleBase):
"""角色响应模型""" """角色响应模型"""
id: int = Field(..., description="角色ID") id: int = Field(..., description="角色ID")
createtime: datetime = Field(..., description="创建时间") createtime: datetime = Field(..., description="创建时间")
updatetime: datetime = Field(..., description="更新时间") updatetime: datetime = Field(..., description="更新时间")
creator: Optional[str] = Field(None, description="创建人") creator: Optional[str] = Field(None, description="创建人")
class Config: class Config:
from_attributes = True from_attributes = True

View File

@ -1,19 +0,0 @@
import urllib.request
# 测试根路径
print("Testing root path...")
try:
response = urllib.request.urlopen('http://127.0.0.1:8000/')
print(f"Status code: {response.status}")
print(f"Response: {response.read().decode('utf-8')}")
except Exception as e:
print(f"Error: {e}")
# 测试/docs路径
print("\nTesting /docs path...")
try:
response = urllib.request.urlopen('http://127.0.0.1:8000/docs')
print(f"Status code: {response.status}")
print(f"Response length: {len(response.read().decode('utf-8'))} characters")
except Exception as e:
print(f"Error: {e}")