[backend] feat: 实现Tasks 4-10的完整代码

- 创建Pydantic Schema模型
- 创建认证中间件和依赖注入
- 创建认证路由
- 创建用户管理路由
- 创建项目管理路由(包含统计功能)
- 创建FastAPI主应用
- 创建环境配置文件
- 修复models的Enum导入问题
- 修复模块导入路径问题

注意:测试仍有部分失败,需要调整错误处理逻辑
This commit is contained in:
xsl
2026-01-26 10:08:14 +08:00
parent c9cd6b9117
commit e7906d57ee
40 changed files with 1213 additions and 17 deletions
BIN
View File
Binary file not shown.
+20
View File
@@ -0,0 +1,20 @@
# 应用配置
APP_NAME=海洋项目管理系统
APP_VERSION=1.0.0
DEBUG=True
# 数据库配置
DB_HOST=localhost
DB_PORT=3306
DB_USER=root
DB_PASSWORD=rootpassword
DB_NAME=project_manager
DB_CHARSET=utf8mb4
# 安全配置
SECRET_KEY=your-secret-key-change-this-in-production
ALGORITHM=HS256
ACCESS_TOKEN_EXPIRE_MINUTES=1440
# CORS配置
CORS_ORIGINS=http://localhost:3000,http://127.0.0.1:3000
Binary file not shown.
Binary file not shown.
Binary file not shown.
+57
View File
@@ -0,0 +1,57 @@
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from config.settings import get_settings
from config.database import init_db
from src.routes import auth_router, users_router, projects_router
settings = get_settings()
app = FastAPI(
title=settings.APP_NAME,
version=settings.APP_VERSION,
description="海洋项目管理系统后端API",
docs_url="/docs",
redoc_url="/redoc",
)
app.add_middleware(
CORSMiddleware,
allow_origins=settings.CORS_ORIGINS,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
app.include_router(auth_router, prefix="/api/v1")
app.include_router(users_router, prefix="/api/v1")
app.include_router(projects_router, prefix="/api/v1")
@app.on_event("startup")
async def startup_event():
await init_db()
@app.get("/")
async def root():
return {
"app": settings.APP_NAME,
"version": settings.APP_VERSION,
"status": "running",
}
@app.get("/health")
async def health():
return {"status": "healthy"}
if __name__ == "__main__":
import uvicorn
uvicorn.run(
"main:app",
host="0.0.0.0",
port=5000,
reload=settings.DEBUG,
)
+1
View File
@@ -3,6 +3,7 @@ testpaths = tests
python_files = test_*.py python_files = test_*.py
python_classes = Test* python_classes = Test*
python_functions = test_* python_functions = test_*
pythonpath = .
addopts = addopts =
-v -v
--strict-markers --strict-markers
Binary file not shown.
+57
View File
@@ -0,0 +1,57 @@
from fastapi import Depends, HTTPException, status
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from config.database import get_db
from src.models.user import User
from src.models.project import Project
from src.middleware.auth import get_current_user
async def get_project_or_404(
project_id: int,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
) -> Project:
result = await db.execute(select(Project).where(Project.id == project_id))
project = result.scalar_one_or_none()
if project is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="项目不存在",
)
return project
async def get_project_if_owner_or_admin(
project_id: int,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
) -> Project:
project = await get_project_or_404(project_id, current_user, db)
if current_user.role != "admin" and project.created_by != current_user.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="无权访问此项目",
)
return project
async def get_user_or_404(
user_id: int,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
) -> User:
result = await db.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if user is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="用户不存在",
)
return user
+13
View File
@@ -0,0 +1,13 @@
from .auth import (
get_current_user,
get_current_active_user,
require_admin,
require_market_or_admin,
)
__all__ = [
"get_current_user",
"get_current_active_user",
"require_admin",
"require_market_or_admin",
]
+81
View File
@@ -0,0 +1,81 @@
from fastapi import Depends, HTTPException, status
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from config.database import get_db
from src.models.user import User
from src.utils.jwt import decode_access_token
security = HTTPBearer()
async def get_current_user(
credentials: HTTPAuthorizationCredentials = Depends(security),
db: AsyncSession = Depends(get_db),
) -> User:
token = credentials.credentials
payload = decode_access_token(token)
if payload is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Token无效或过期",
headers={"WWW-Authenticate": "Bearer"},
)
user_id = payload.get("sub")
if user_id is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Token无效",
)
result = await db.execute(select(User).where(User.id == int(user_id)))
user = result.scalar_one_or_none()
if user is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户不存在",
)
if not user.is_active:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户已被禁用",
)
return user
async def get_current_active_user(
current_user: User = Depends(get_current_user),
) -> User:
if not current_user.is_active:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="用户未激活",
)
return current_user
async def require_admin(
current_user: User = Depends(get_current_active_user),
) -> User:
if current_user.role != "admin":
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="需要管理员权限",
)
return current_user
async def require_market_or_admin(
current_user: User = Depends(get_current_active_user),
) -> User:
if current_user.role not in ["admin", "market"]:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="需要市场部或管理员权限",
)
return current_user
Binary file not shown.
+2 -7
View File
@@ -2,12 +2,7 @@
from sqlalchemy import Column, Integer, String, Text, Date, Numeric, DateTime, Enum, ForeignKey from sqlalchemy import Column, Integer, String, Text, Date, Numeric, DateTime, Enum, ForeignKey
from sqlalchemy.sql import func from sqlalchemy.sql import func
from sqlalchemy.orm import relationship from sqlalchemy.orm import relationship
from ..config.database import Base from config.database import Base
class AdjustmentEnum(str, Enum):
YES = "yes"
NO = "no"
class Project(Base): class Project(Base):
@@ -48,7 +43,7 @@ class Project(Base):
# 成本控制 # 成本控制
total_cost_control = Column(Numeric(15, 2), comment="总体成本(控制)") total_cost_control = Column(Numeric(15, 2), comment="总体成本(控制)")
is_adjusted = Column(Enum(AdjustmentEnum), default=AdjustmentEnum.NO, index=True, comment="是否调整") is_adjusted = Column(Enum("", "", name="is_adjusted"), default="", index=True, comment="是否调整")
labor_cost_control = Column(Numeric(15, 2), comment="人工成本(控制)") labor_cost_control = Column(Numeric(15, 2), comment="人工成本(控制)")
labor_cost_planned = Column(Numeric(15, 2), comment="农民工工资(计划)") labor_cost_planned = Column(Numeric(15, 2), comment="农民工工资(计划)")
labor_cost_paid = Column(Numeric(15, 2), comment="农民工工资(实付)") labor_cost_paid = Column(Numeric(15, 2), comment="农民工工资(实付)")
+2 -8
View File
@@ -2,13 +2,7 @@
from sqlalchemy import Column, Integer, String, Boolean, DateTime, Enum, ForeignKey from sqlalchemy import Column, Integer, String, Boolean, DateTime, Enum, ForeignKey
from sqlalchemy.sql import func from sqlalchemy.sql import func
from sqlalchemy.orm import relationship from sqlalchemy.orm import relationship
from ..config.database import Base from config.database import Base
class UserRoleEnum(str, Enum):
admin = "admin"
market = "market"
other = "other"
class User(Base): class User(Base):
@@ -19,7 +13,7 @@ class User(Base):
password_hash = Column(String(255), nullable=False, comment="密码哈希") password_hash = Column(String(255), nullable=False, comment="密码哈希")
real_name = Column(String(100), nullable=False, comment="真实姓名") real_name = Column(String(100), nullable=False, comment="真实姓名")
department = Column(String(50), nullable=False, index=True, comment="部门") department = Column(String(50), nullable=False, index=True, comment="部门")
role = Column(Enum(UserRoleEnum), nullable=False, index=True, comment="角色") role = Column(Enum("admin", "market", "other", name="user_role"), nullable=False, index=True, comment="角色")
email = Column(String(100), unique=True, index=True, comment="邮箱") email = Column(String(100), unique=True, index=True, comment="邮箱")
phone = Column(String(20), comment="电话") phone = Column(String(20), comment="电话")
is_active = Column(Boolean, default=True, index=True, comment="是否激活") is_active = Column(Boolean, default=True, index=True, comment="是否激活")
+9
View File
@@ -0,0 +1,9 @@
from .auth import router as auth_router
from .users import router as users_router
from .projects import router as projects_router
__all__ = [
"auth_router",
"users_router",
"projects_router",
]
Binary file not shown.
Binary file not shown.
+82
View File
@@ -0,0 +1,82 @@
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from config.database import get_db
from src.models.user import User
from src.schemas.auth import LoginRequest, LoginResponse
from src.schemas.user import UserResponse
from src.utils.password import verify_password
from src.utils.jwt import create_access_token
from src.middleware.auth import get_current_user
router = APIRouter(prefix="/auth", tags=["认证"])
@router.post("/login", response_model=dict)
async def login(
login_data: LoginRequest,
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(User).where(User.username == login_data.username)
)
user = result.scalar_one_or_none()
if not user or not verify_password(login_data.password, user.password_hash):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail={
"success": False,
"message": "用户名或密码错误",
"data": None,
"error_code": "1002",
}
)
if not user.is_active:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail={
"success": False,
"message": "用户未激活",
"data": None,
"error_code": "1002",
}
)
access_token = create_access_token(
data={"sub": str(user.id), "username": user.username}
)
return {
"success": True,
"message": "登录成功",
"data": {
"token": access_token,
"token_type": "bearer",
"user": UserResponse.model_validate(user),
},
"error_code": None,
}
@router.get("/me", response_model=dict)
async def get_current_user_info(
current_user: User = Depends(get_current_user),
):
return {
"success": True,
"message": "获取成功",
"data": UserResponse.model_validate(current_user),
"error_code": None,
}
@router.post("/logout", response_model=dict)
async def logout(current_user: User = Depends(get_current_user)):
return {
"success": True,
"message": "登出成功",
"data": None,
"error_code": None,
}
+432
View File
@@ -0,0 +1,432 @@
from fastapi import APIRouter, Depends, Query
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, func, and_, desc, asc
from config.database import get_db
from src.models.user import User
from src.models.project import Project
from src.schemas.project import (
ProjectCreate,
ProjectUpdate,
ProjectResponse,
ProjectListResponse,
)
from src.middleware.auth import (
get_current_user,
require_market_or_admin,
)
from src.dependencies import get_project_or_404
from typing import Optional
from datetime import date
from decimal import Decimal
router = APIRouter(prefix="/projects", tags=["项目管理"])
@router.get("", response_model=dict)
async def get_projects(
page: int = Query(1, ge=1),
page_size: int = Query(10, ge=1, le=100),
project_no: Optional[str] = None,
engineering_type: Optional[str] = None,
project_department: Optional[str] = None,
signing_date_start: Optional[date] = None,
signing_date_end: Optional[date] = None,
contract_amount_min: Optional[Decimal] = None,
contract_amount_max: Optional[Decimal] = None,
keyword: Optional[str] = None,
sort_by: Optional[str] = None,
sort_order: Optional[str] = Query(None, pattern="^(asc|desc)$"),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
query = select(Project)
conditions = []
if project_no:
conditions.append(Project.project_no == project_no)
if engineering_type:
conditions.append(Project.engineering_type == engineering_type)
if project_department:
conditions.append(Project.project_department == project_department)
if signing_date_start:
conditions.append(Project.signing_date >= signing_date_start)
if signing_date_end:
conditions.append(Project.signing_date <= signing_date_end)
if contract_amount_min:
conditions.append(Project.contract_amount >= contract_amount_min)
if contract_amount_max:
conditions.append(Project.contract_amount <= contract_amount_max)
if keyword:
conditions.append(
(Project.name.contains(keyword)) |
(Project.owner_unit.contains(keyword))
)
if conditions:
query = query.where(and_(*conditions))
if sort_by:
sort_column = getattr(Project, sort_by, None)
if sort_column:
order = desc if sort_order == "desc" else asc
query = query.order_by(order(sort_column))
else:
query = query.order_by(Project.created_at.desc())
total_query = select(func.count()).select_from(query.subquery())
total_result = await db.execute(total_query)
total = total_result.scalar()
query = query.offset((page - 1) * page_size).limit(page_size)
result = await db.execute(query)
projects = result.scalars().all()
project_data = []
for p in projects:
creator_result = await db.execute(
select(User.real_name).where(User.id == p.created_by)
)
creator_name = creator_result.scalar_one_or_none()
p_dict = ProjectResponse.model_validate(p).model_dump()
p_dict["created_by_name"] = creator_name
project_data.append(p_dict)
return {
"success": True,
"message": "获取成功",
"data": {
"items": project_data,
"total": total,
"page": page,
"page_size": page_size,
},
"error_code": None,
}
@router.get("/{project_id}", response_model=dict)
async def get_project(
project: Project = Depends(get_project_or_404),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(User.real_name).where(User.id == project.created_by)
)
creator_name = result.scalar_one_or_none()
project_dict = ProjectResponse.model_validate(project).model_dump()
project_dict["created_by_name"] = creator_name
return {
"success": True,
"message": "获取成功",
"data": project_dict,
"error_code": None,
}
@router.post("", response_model=dict)
async def create_project(
project_data: ProjectCreate,
current_user: User = Depends(require_market_or_admin),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(Project).where(Project.project_no == project_data.project_no)
)
if result.scalar_one_or_none():
return {
"success": False,
"message": "项目编号已存在",
"data": None,
"error_code": "2002",
}
project = Project(
**project_data.model_dump(),
created_by=current_user.id,
)
db.add(project)
await db.commit()
await db.refresh(project)
return {
"success": True,
"message": "项目创建成功",
"data": ProjectResponse.model_validate(project),
"error_code": None,
}
@router.put("/{project_id}", response_model=dict)
async def update_project(
project_id: int,
project_data: ProjectUpdate,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(Project).where(Project.id == project_id))
project = result.scalar_one_or_none()
if not project:
return {
"success": False,
"message": "项目不存在",
"data": None,
"error_code": "2001",
}
if current_user.role == "market" and project.created_by != current_user.id:
return {
"success": False,
"message": "无权修改此项目",
"data": None,
"error_code": "3001",
}
update_data = project_data.model_dump(exclude_unset=True)
for field, value in update_data.items():
setattr(project, field, value)
await db.commit()
await db.refresh(project)
return {
"success": True,
"message": "项目更新成功",
"data": ProjectResponse.model_validate(project),
"error_code": None,
}
@router.delete("/{project_id}", response_model=dict)
async def delete_project(
project_id: int,
current_user: User = Depends(require_market_or_admin),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(Project).where(Project.id == project_id))
project = result.scalar_one_or_none()
if not project:
return {
"success": False,
"message": "项目不存在",
"data": None,
"error_code": "2001",
}
if current_user.role == "market" and project.created_by != current_user.id:
return {
"success": False,
"message": "无权删除此项目",
"data": None,
"error_code": "3001",
}
await db.delete(project)
await db.commit()
return {
"success": True,
"message": "项目删除成功",
"data": None,
"error_code": None,
}
@router.get("/statistics/basic", response_model=dict)
async def get_statistics(
engineering_type: Optional[str] = None,
signing_date_start: Optional[date] = None,
signing_date_end: Optional[date] = None,
contract_amount_min: Optional[Decimal] = None,
contract_amount_max: Optional[Decimal] = None,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
query = select(Project)
conditions = []
if engineering_type:
conditions.append(Project.engineering_type == engineering_type)
if signing_date_start:
conditions.append(Project.signing_date >= signing_date_start)
if signing_date_end:
conditions.append(Project.signing_date <= signing_date_end)
if contract_amount_min:
conditions.append(Project.contract_amount >= contract_amount_min)
if contract_amount_max:
conditions.append(Project.contract_amount <= contract_amount_max)
if conditions:
query = query.where(and_(*conditions))
result = await db.execute(query)
projects = result.scalars().all()
total_count = len(projects)
total_contract_amount = sum(
p.contract_amount or Decimal(0) for p in projects
)
total_receipt_amount = sum(
p.actual_receipt_amount or Decimal(0) for p in projects
)
total_payment_amount = sum(
p.actual_payment_amount or Decimal(0) for p in projects
)
avg_receipt_rate = 0
if total_contract_amount > 0:
avg_receipt_rate = (
(total_receipt_amount / total_contract_amount) * 100
)
avg_payment_rate = 0
total_payable = sum(p.payable_amount or Decimal(0) for p in projects)
if total_payable > 0:
avg_payment_rate = (
(total_payment_amount / total_payable) * 100
)
avg_progress = 0
if total_count > 0:
avg_progress = sum(
p.cumulative_progress or Decimal(0) for p in projects
) / total_count
return {
"success": True,
"message": "统计成功",
"data": {
"total_count": total_count,
"total_contract_amount": float(total_contract_amount),
"total_receipt_amount": float(total_receipt_amount),
"total_payment_amount": float(total_payment_amount),
"avg_receipt_completion_rate": float(avg_receipt_rate),
"avg_payment_completion_rate": float(avg_payment_rate),
"avg_cumulative_progress": float(avg_progress),
},
"error_code": None,
}
@router.get("/statistics/group", response_model=dict)
async def get_group_statistics(
group_by: str = Query(..., pattern="^(engineering_type|project_department)$"),
signing_date_start: Optional[date] = None,
signing_date_end: Optional[date] = None,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
query = select(Project)
if signing_date_start:
query = query.where(Project.signing_date >= signing_date_start)
if signing_date_end:
query = query.where(Project.signing_date <= signing_date_end)
result = await db.execute(query)
projects = result.scalars().all()
groups = {}
for p in projects:
group_value = getattr(p, group_by)
if group_value not in groups:
groups[group_value] = {
"count": 0,
"total_contract_amount": Decimal(0),
"total_receipt_amount": Decimal(0),
"total_payment_amount": Decimal(0),
}
groups[group_value]["count"] += 1
groups[group_value]["total_contract_amount"] += p.contract_amount or Decimal(0)
groups[group_value]["total_receipt_amount"] += p.actual_receipt_amount or Decimal(0)
groups[group_value]["total_payment_amount"] += p.actual_payment_amount or Decimal(0)
result_list = []
for key, value in groups.items():
if group_by == "engineering_type":
item = {"engineering_type": key}
else:
item = {"project_department": key}
item.update({
"count": value["count"],
"total_contract_amount": float(value["total_contract_amount"]),
"total_receipt_amount": float(value["total_receipt_amount"]),
"total_payment_amount": float(value["total_payment_amount"]),
})
result_list.append(item)
return {
"success": True,
"message": "统计成功",
"data": result_list,
"error_code": None,
}
@router.get("/statistics/timeline", response_model=dict)
async def get_timeline_statistics(
time_field: str = Query(..., pattern="^(signing_date|start_date|planned_end_date)$"),
group_by: str = Query("month", pattern="^(day|month|year)$"),
engineering_type: Optional[str] = None,
start_date: Optional[date] = None,
end_date: Optional[date] = None,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
query = select(Project)
if engineering_type:
query = query.where(Project.engineering_type == engineering_type)
if start_date:
date_column = getattr(Project, time_field)
query = query.where(date_column >= start_date)
if end_date:
date_column = getattr(Project, time_field)
query = query.where(date_column <= end_date)
result = await db.execute(query)
projects = result.scalars().all()
groups = {}
for p in projects:
date_value = getattr(p, time_field)
if not date_value:
continue
if group_by == "day":
key = date_value.strftime("%Y-%m-%d")
elif group_by == "month":
key = date_value.strftime("%Y-%m")
else:
key = str(date_value.year)
if key not in groups:
groups[key] = {"count": 0, "total_contract_amount": Decimal(0)}
groups[key]["count"] += 1
groups[key]["total_contract_amount"] += p.contract_amount or Decimal(0)
result_list = []
for key in sorted(groups.keys()):
value = groups[key]
if group_by == "day":
item = {"day": key}
elif group_by == "month":
item = {"month": key}
else:
item = {"year": key}
item.update({
"count": value["count"],
"total_contract_amount": float(value["total_contract_amount"]),
})
result_list.append(item)
return {
"success": True,
"message": "统计成功",
"data": result_list,
"error_code": None,
}
+228
View File
@@ -0,0 +1,228 @@
from fastapi import APIRouter, Depends, Query
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, func, or_
from config.database import get_db
from src.models.user import User
from src.schemas.user import (
UserCreate,
UserUpdate,
UserResponse,
UserListResponse,
PasswordReset,
)
from src.utils.password import hash_password
from src.middleware.auth import require_admin
from typing import Optional
router = APIRouter(prefix="/users", tags=["用户管理"])
@router.get("", response_model=dict)
async def get_users(
page: int = Query(1, ge=1),
page_size: int = Query(10, ge=1, le=100),
department: Optional[str] = None,
role: Optional[str] = None,
keyword: Optional[str] = None,
current_user: User = Depends(require_admin),
db: AsyncSession = Depends(get_db),
):
query = select(User)
if department:
query = query.where(User.department == department)
if role:
query = query.where(User.role == role)
if keyword:
query = query.where(
or_(
User.username.contains(keyword),
User.real_name.contains(keyword),
User.email.contains(keyword),
)
)
total_query = select(func.count()).select_from(query.subquery())
total_result = await db.execute(total_query)
total = total_result.scalar()
query = query.offset((page - 1) * page_size).limit(page_size)
result = await db.execute(query)
users = result.scalars().all()
return {
"success": True,
"message": "获取成功",
"data": {
"items": [UserResponse.model_validate(u) for u in users],
"total": total,
"page": page,
"page_size": page_size,
},
"error_code": None,
}
@router.post("", response_model=dict)
async def create_user(
user_data: UserCreate,
current_user: User = Depends(require_admin),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(User).where(User.username == user_data.username)
)
if result.scalar_one_or_none():
return {
"success": False,
"message": "用户名已存在",
"data": None,
"error_code": "2002",
}
if user_data.email:
result = await db.execute(
select(User).where(User.email == user_data.email)
)
if result.scalar_one_or_none():
return {
"success": False,
"message": "邮箱已存在",
"data": None,
"error_code": "2002",
}
user = User(
username=user_data.username,
password_hash=hash_password(user_data.password),
real_name=user_data.real_name,
department=user_data.department,
role=user_data.role,
email=user_data.email,
phone=user_data.phone,
)
db.add(user)
await db.commit()
await db.refresh(user)
return {
"success": True,
"message": "用户创建成功",
"data": UserResponse.model_validate(user),
"error_code": None,
}
@router.get("/{user_id}", response_model=dict)
async def get_user(
user_id: int,
current_user: User = Depends(require_admin),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if not user:
return {
"success": False,
"message": "用户不存在",
"data": None,
"error_code": "2001",
}
return {
"success": True,
"message": "获取成功",
"data": UserResponse.model_validate(user),
"error_code": None,
}
@router.put("/{user_id}", response_model=dict)
async def update_user(
user_id: int,
user_data: UserUpdate,
current_user: User = Depends(require_admin),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if not user:
return {
"success": False,
"message": "用户不存在",
"data": None,
"error_code": "2001",
}
update_data = user_data.model_dump(exclude_unset=True)
for field, value in update_data.items():
setattr(user, field, value)
await db.commit()
await db.refresh(user)
return {
"success": True,
"message": "用户更新成功",
"data": UserResponse.model_validate(user),
"error_code": None,
}
@router.delete("/{user_id}", response_model=dict)
async def delete_user(
user_id: int,
current_user: User = Depends(require_admin),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if not user:
return {
"success": False,
"message": "用户不存在",
"data": None,
"error_code": "2001",
}
await db.delete(user)
await db.commit()
return {
"success": True,
"message": "用户删除成功",
"data": None,
"error_code": None,
}
@router.post("/{user_id}/reset-password", response_model=dict)
async def reset_password(
user_id: int,
password_data: PasswordReset,
current_user: User = Depends(require_admin),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if not user:
return {
"success": False,
"message": "用户不存在",
"data": None,
"error_code": "2001",
}
user.password_hash = hash_password(password_data.new_password)
await db.commit()
return {
"success": True,
"message": "密码重置成功",
"data": None,
"error_code": None,
}
+28
View File
@@ -0,0 +1,28 @@
from .auth import LoginRequest, LoginResponse
from .user import (
UserCreate,
UserUpdate,
UserResponse,
UserListResponse,
PasswordReset,
)
from .project import (
ProjectCreate,
ProjectUpdate,
ProjectResponse,
ProjectListResponse,
)
__all__ = [
"LoginRequest",
"LoginResponse",
"UserCreate",
"UserUpdate",
"UserResponse",
"UserListResponse",
"PasswordReset",
"ProjectCreate",
"ProjectUpdate",
"ProjectResponse",
"ProjectListResponse",
]
Binary file not shown.
Binary file not shown.
+14
View File
@@ -0,0 +1,14 @@
from pydantic import BaseModel, Field
from typing import Optional
from .user import UserResponse
class LoginRequest(BaseModel):
username: str = Field(..., min_length=3)
password: str = Field(..., min_length=6)
class LoginResponse(BaseModel):
token: str
token_type: str = "bearer"
user: UserResponse
+136
View File
@@ -0,0 +1,136 @@
from pydantic import BaseModel, Field
from typing import Optional
from datetime import date, datetime
from decimal import Decimal
class ProjectBase(BaseModel):
project_no: str = Field(..., max_length=50)
name: str = Field(..., max_length=200)
engineering_type: str = Field(..., max_length=50)
contract_amount: Optional[Decimal] = Field(None, ge=0)
class ProjectCreate(ProjectBase):
power_contract_no: Optional[str] = Field(None, max_length=100)
subitem_count: Optional[int] = Field(None, ge=0)
subitem_code: Optional[str] = Field(None, max_length=50)
total_investment: Optional[Decimal] = Field(None, ge=0)
warranty_ratio: Optional[Decimal] = None
settlement_amount: Optional[Decimal] = None
total_cost_estimated: Optional[Decimal] = None
voltage_level: Optional[str] = Field(None, max_length=50)
owner_unit: Optional[str] = Field(None, max_length=200)
owner_contact: Optional[str] = Field(None, max_length=200)
bidding_type: Optional[str] = Field(None, max_length=50)
signing_date: Optional[date] = None
start_date: Optional[date] = None
planned_end_date: Optional[date] = None
project_department: Optional[str] = Field(None, max_length=100)
project_leader: Optional[str] = Field(None, max_length=200)
payment_method: Optional[str] = None
total_cost_control: Optional[Decimal] = None
labor_cost_control: Optional[Decimal] = None
material_cost_control: Optional[Decimal] = None
other_cost_control: Optional[Decimal] = None
remarks: Optional[str] = None
class ProjectUpdate(BaseModel):
name: Optional[str] = Field(None, max_length=200)
engineering_type: Optional[str] = Field(None, max_length=50)
contract_amount: Optional[Decimal] = Field(None, ge=0)
total_investment: Optional[Decimal] = None
settlement_amount: Optional[Decimal] = None
actual_receipt_amount: Optional[Decimal] = None
actual_payment_amount: Optional[Decimal] = None
invoice_amount: Optional[Decimal] = None
payable_amount: Optional[Decimal] = None
cumulative_progress: Optional[Decimal] = None
labor_cost_paid: Optional[Decimal] = None
material_cost_actual: Optional[Decimal] = None
material_cost_paid: Optional[Decimal] = None
other_cost_actual: Optional[Decimal] = None
tax_amount: Optional[Decimal] = None
profit: Optional[Decimal] = None
remarks: Optional[str] = None
problems: Optional[str] = None
suggestions: Optional[str] = None
class ProjectResponse(BaseModel):
id: int
project_no: str
power_contract_no: Optional[str]
name: str
subitem_count: Optional[int]
subitem_code: Optional[str]
total_investment: Optional[Decimal]
contract_amount: Optional[Decimal]
warranty_ratio: Optional[Decimal]
settlement_amount: Optional[Decimal]
total_cost_estimated: Optional[Decimal]
voltage_level: Optional[str]
engineering_type: Optional[str]
owner_unit: Optional[str]
owner_contact: Optional[str]
bidding_type: Optional[str]
signing_date: Optional[date]
start_date: Optional[date]
planned_end_date: Optional[date]
actual_end_date: Optional[date]
warranty_amount: Optional[Decimal]
warranty_expiry_date: Optional[date]
actual_warranty_refund_date: Optional[date]
project_department: Optional[str]
project_leader: Optional[str]
payment_method: Optional[str]
total_cost_control: Optional[Decimal]
is_adjusted: Optional[str]
labor_cost_control: Optional[Decimal]
labor_cost_planned: Optional[Decimal]
labor_cost_paid: Optional[Decimal]
material_cost_control: Optional[Decimal]
material_cost_payable: Optional[Decimal]
material_cost_actual: Optional[Decimal]
material_cost_paid: Optional[Decimal]
other_cost_control: Optional[Decimal]
other_cost_payable: Optional[Decimal]
other_cost_actual: Optional[Decimal]
tax_amount: Optional[Decimal]
profit: Optional[Decimal]
actual_profit: Optional[Decimal]
cost_settlement_amount: Optional[Decimal]
cumulative_progress: Optional[Decimal]
receivable_amount: Optional[Decimal]
invoice_amount: Optional[Decimal]
actual_receipt_amount: Optional[Decimal]
receipt_completion_rate: Optional[Decimal]
payable_amount: Optional[Decimal]
actual_payment_amount: Optional[Decimal]
unpaid_amount: Optional[Decimal]
payment_completion_rate: Optional[Decimal]
labor_debt_amount: Optional[Decimal]
settlement_cost_amount: Optional[Decimal]
settlement_labor_cost: Optional[Decimal]
settlement_material_cost: Optional[Decimal]
settlement_other_cost: Optional[Decimal]
due_settlement_count: Optional[int]
unsettlement_count: Optional[int]
problems: Optional[str]
suggestions: Optional[str]
remarks: Optional[str]
created_by: int
created_by_name: Optional[str] = None
created_at: datetime
updated_at: Optional[datetime]
class Config:
from_attributes = True
class ProjectListResponse(BaseModel):
items: list[ProjectResponse]
total: int
page: int
page_size: int
+44
View File
@@ -0,0 +1,44 @@
from pydantic import BaseModel, EmailStr, Field
from typing import Optional
from datetime import datetime
class UserBase(BaseModel):
username: str = Field(..., min_length=3, max_length=50)
real_name: str = Field(..., max_length=100)
department: str = Field(..., max_length=50)
role: str = Field(..., pattern="^(admin|market|other)$")
email: Optional[EmailStr] = None
phone: Optional[str] = Field(None, max_length=20)
class UserCreate(UserBase):
password: str = Field(..., min_length=6)
class UserUpdate(BaseModel):
real_name: Optional[str] = Field(None, max_length=100)
email: Optional[EmailStr] = None
phone: Optional[str] = Field(None, max_length=20)
is_active: Optional[bool] = None
class PasswordReset(BaseModel):
new_password: str = Field(..., min_length=6)
class UserResponse(UserBase):
id: int
is_active: bool
created_at: datetime
updated_at: Optional[datetime]
class Config:
from_attributes = True
class UserListResponse(BaseModel):
items: list[UserResponse]
total: int
page: int
page_size: int
Binary file not shown.
+7 -2
View File
@@ -1,3 +1,8 @@
import sys
from pathlib import Path
backend_dir = Path(__file__).parent.parent
sys.path.insert(0, str(backend_dir))
import pytest import pytest
import asyncio import asyncio
from typing import AsyncGenerator, Generator from typing import AsyncGenerator, Generator
@@ -5,8 +10,8 @@ from httpx import AsyncClient
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker
from sqlalchemy.pool import StaticPool from sqlalchemy.pool import StaticPool
from src.main import app from main import app
from src.config.database import Base, get_db from config.database import Base, get_db
from src.models.user import User from src.models.user import User
from src.utils.password import hash_password from src.utils.password import hash_password