欢迎光临
我们一直在努力

零基础入门python24:分类 CRUD 与用户数据隔离

零基础入门python24:分类 CRUD 与用户数据隔离

Flask 项目分层图

一、上一篇课后练习讲解

上一篇要求配置会话过期和 Cookie 安全属性。开发环境可以使用较短过期时间方便测试,生产环境要启用 SESSION_COOKIE_SECURE、HTTPOnly 和合适的 SameSite 策略。

上一篇课后练习完整答案

上一篇练习的要求已落实到下面完整文件;先运行项目测试,再用 curl 对照状态码和数据库持久化结果
答案要点:登录用 check_password_hash,user_loader 从数据库加载,/me 依赖 current_user;未登录统一 401,不暴露邮箱是否存在。

文件:app/auth.py

完整参考答案文件

完整文件:app/auth.py

from flask import Blueprint, request
from flask_login import login_user, login_required, current_user
from werkzeug.security import check_password_hash
from .models import User
bp = Blueprint("auth", __name__, url_prefix="/api/auth")
@bp.post("/login")
def login():
data = request.get_json(silent=True) or {}
user = User.query.filter_by(email=str(data.get("email", "")).casefold()).first()
if not user or not check_password_hash(user.password_hash, str(data.get("password", ""))):
return {"error": "invalid_credentials"}, 401
login_user(user)
return {"id": user.id, "email": user.email}
@bp.get("/me")
@login_required
def me():
return {"id": current_user.id, "email": current_user.email}

完整参考答案文件

本篇对应的交付源码完整文件:flask-ledger/app/auth.py

from flask import Blueprint, request
from flask_login import current_user, login_required, login_user, logout_user

from .extensions import db
from .models import User

bp = Blueprint("auth", __name__, url_prefix="/api/auth")

@bp.post("/register")
def register():
data = request.get_json(silent=True) or {}
email = str(data.get("email", "")).strip().lower()
password = str(data.get("password", ""))
if "@" not in email:
return {"message": "邮箱格式不正确"}, 400
if len(password) < 8:
return {"message": "密码至少8位"}, 400
if db.session.scalar(db.select(User).where(User.email == email)):
return {"message": "邮箱已注册"}, 409
user = User(email=email)
user.set_password(password)
db.session.add(user)
db.session.commit()
return {"id": user.id, "email": user.email}, 201

@bp.post("/login")
def login():
data = request.get_json(silent=True) or {}
email = str(data.get("email", "")).strip().lower()
user = db.session.scalar(db.select(User).where(User.email == email))
if not user or not user.check_password(str(data.get("password", ""))):
return {"message": "邮箱或密码错误"}, 401
login_user(user)
return {"id": user.id, "email": user.email}

@bp.post("/logout")
@login_required
def logout():
logout_user()
return {"message": "已退出"}

@bp.get("/me")
@login_required
def me():
return {"id": current_user.id, "email": current_user.email}

验收命令:python -m pytest -q(Django 项目使用 python manage.py test)。预期测试通过;若失败先检查迁移、配置和事务回滚。

二、本篇完成什么

分类是账目的前置资源。本篇实现当前用户的分类新增、列表、修改和删除,并验证用户 A 不能读取或修改用户 B 的分类。

本篇流程图

三、为什么每个查询都要带 user_id

只用 Category.query.get(category_id) 会让用户猜 id 后读到别人的数据。正确做法是把资源所有权放进查询条件:where(Category.id == id, Category.user_id == current_user.id)。这样即使客户端传入别人的 id,查询也不会返回记录。

四、关键实现

@bp.post('/api/categories')
@login_required
def create_category():
data = request.get_json(silent=True) or {}
name = str(data.get('name', '')).strip()
if not name:
return {'message': '分类名不能为空'}, 400
if db.session.scalar(db.select(Category).where(
Category.user_id == current_user.id, Category.name == name
)):
return {'message': '分类已存在'}, 409
row = Category(name=name, user_id=current_user.id)
db.session.add(row)
db.session.commit()
return {'id': row.id, 'name': row.name}, 201

重复检查是用户范围内的检查,数据库联合唯一约束是最终防线;两者同时存在,前者提供友好错误,后者应对并发请求。

五、验收

用户 A 登录后创建“餐饮”,列表能看到;用户 B 登录后列表为空,修改 A 的分类返回 404。测试同时覆盖未登录 401、空名称 400、重复名称 409。课后练习:增加分类删除接口,并阻止删除仍有账目的分类。

项目增量:分类 CRUD 与用户隔离

分类也是用户资源。查询、修改和删除都必须带 current_user.id;只根据分类 id 查询会导致越权读取。

category = Category.query.filter_by(id=category_id, user_id=current_user.id).first_or_404()
category.name = form.name.strip()
db.session.commit()

名称唯一约束应按 user_id + name 组合,而不是全站唯一。删除仍被账目使用的分类要返回业务错误或要求迁移账目,不能静默丢失历史。

验收与课后练习

用户 A 不能读取、修改或删除用户 B 的分类;重复名称返回 409;课后增加分类删除保护测试。

五、分类 CRUD 的完整闭环

分类接口不是简单的四个 SQL。每一步都要带上当前用户条件,避免用户 A 通过猜 id 访问用户 B 的数据:

@ledger_bp.get("/categories")
@login_required
def list_categories():
rows = (Category.query.filter_by(user_id=current_user.id)
.order_by(Category.name.asc()).all())
return {"items": [{"id": row.id, "name": row.name} for row in rows]}

@ledger_bp.post("/categories")
@login_required
def create_category():
name = (request.get_json(silent=True) or {}).get("name", "").strip()
if not 1 <= len(name) <= 30:
return {"error": "name_length"}, 400
row = Category(user_id=current_user.id, name=name)
db.session.add(row)
try:
db.session.commit()
except IntegrityError:
db.session.rollback()
return {"error": "category_exists"}, 409
return {"id": row.id, "name": row.name}, 201

@ledger_bp.delete("/categories/<int:category_id>")
@login_required
def delete_category(category_id):
row = Category.query.filter_by(id=category_id, user_id=current_user.id).first()
if row is None:
# 对不存在和不属于自己的资源统一返回 404,避免泄露资源存在性。
return {"error": "not_found"}, 404
db.session.delete(row)
db.session.commit()
return "", 204

删除分类前要决定已有账目如何处理:本项目把 category_id 设为可空,删除时置空;如果产品要求禁止删除,就应该在服务层先统计引用数量并返回 409。不要依赖 SQLite 的默认外键行为,因为不同数据库的外键开关不同。

六、分页和搜索的边界

可以用 page、size 两个参数,并限制 size <= 100:

page = max(request.args.get("page", 1, type=int), 1)
size = min(max(request.args.get("size", 20, type=int), 1), 100)
keyword = request.args.get("q", "").strip()
query = Category.query.filter_by(user_id=current_user.id)
if keyword:
query = query.filter(Category.name.ilike(f"%{keyword}%"))
total = query.count()
items = query.order_by(Category.id.desc()).offset((page1)*size).limit(size).all()

count() 和列表查询是两次 SQL,数据量很大时可以改为游标分页;现在先保证契约清晰。不要直接把 size 拼到 SQL 字符串里,ORM 参数化会自动处理类型。

七、上一篇练习讲解与错误排查

上一篇的 /auth/me 要复用 current_user,不能从请求体接收 user id。测试至少覆盖:A 创建分类后能看到,B 看不到 A 的分类;重复名称返回 409;删除不存在的 id 返回 404。IntegrityError 未回滚会让后续请求一直处于“事务已失败”状态,表现为 PendingRollbackError,遇到该错误第一时间执行 db.session.rollback()。

八、本篇练习

实现分类更新 PUT /categories/<id>,要求新名称仍在 1—30 字且不能与当前用户其他分类重复;补充一个“引用账目时拒绝删除”的模式开关。下一篇将把分类 id 接入账目创建,并讲 Decimal 与日期输入。

本篇结束:完整模块文件

下面是交付项目中真实存在的完整文件 flask-ledger/app/ledger.py。它覆盖本篇新增逻辑以及前文已经完成的依赖代码;复制单个函数会丢失上下文,因此这里提供整份文件。

from datetime import date
from decimal import Decimal, InvalidOperation

from flask import Blueprint, request
from flask_login import current_user, login_required
from sqlalchemy import func

from .extensions import db
from .models import Category, Transaction

bp = Blueprint("ledger", __name__, url_prefix="/api")

def owned_category(category_id: int):
return db.session.scalar(
db.select(Category).where(Category.id == category_id, Category.user_id == current_user.id)
)

@bp.get("/categories")
@login_required
def list_categories():
rows = db.session.scalars(
db.select(Category).where(Category.user_id == current_user.id).order_by(Category.name)
).all()
return [{"id": row.id, "name": row.name} for row in rows]

@bp.post("/categories")
@login_required
def create_category():
name = str((request.get_json(silent=True) or {}).get("name", "")).strip()
if not 1 <= len(name) <= 40:
return {"message": "分类名称长度应为1到40"}, 400
exists = db.session.scalar(
db.select(Category).where(Category.user_id == current_user.id, Category.name == name)
)
if exists:
return {"message": "分类已存在"}, 409
row = Category(name=name, user_id=current_user.id)
db.session.add(row)
db.session.commit()
return {"id": row.id, "name": row.name}, 201

@bp.post("/transactions")
@login_required
def create_transaction():
data = request.get_json(silent=True) or {}
try:
amount = Decimal(str(data.get("amount", "0"))).quantize(Decimal("0.01"))
happened_on = date.fromisoformat(str(data.get("happened_on", date.today())))
category_id = int(data.get("category_id"))
except (InvalidOperation, ValueError, TypeError):
return {"message": "金额、日期或分类格式不正确"}, 400
kind = str(data.get("kind", ""))
if kind not in {"income", "expense"} or amount <= 0:
return {"message": "类型或金额不正确"}, 400
if not owned_category(category_id):
return {"message": "分类不存在"}, 404
row = Transaction(
kind=kind, amount=amount, note=str(data.get("note", ""))[:200],
happened_on=happened_on, category_id=category_id, user_id=current_user.id,
)
db.session.add(row)
db.session.commit()
return row.to_dict(), 201

@bp.get("/transactions")
@login_required
def list_transactions():
page = max(request.args.get("page", 1, type=int), 1)
size = min(max(request.args.get("size", 10, type=int), 1), 100)
query = db.select(Transaction).where(Transaction.user_id == current_user.id)
if kind := request.args.get("kind"):
query = query.where(Transaction.kind == kind)
if category_id := request.args.get("category_id", type=int):
query = query.where(Transaction.category_id == category_id)
rows = db.session.scalars(
query.order_by(Transaction.happened_on.desc(), Transaction.id.desc())
.offset((page 1) * size).limit(size)
).all()
return {"page": page, "size": size, "items": [row.to_dict() for row in rows]}

@bp.delete("/transactions/<int:transaction_id>")
@login_required
def delete_transaction(transaction_id: int):
row = db.session.scalar(
db.select(Transaction).where(
Transaction.id == transaction_id, Transaction.user_id == current_user.id
)
)
if not row:
return {"message": "账目不存在"}, 404
db.session.delete(row)
db.session.commit()
return "", 204

@bp.get("/statistics/monthly")
@login_required
def monthly_statistics():
month = request.args.get("month", date.today().strftime("%Y-%m"))
try:
start = date.fromisoformat(month + "-01")
except ValueError:
return {"message": "月份格式应为YYYY-MM"}, 400
end = date(start.year + (start.month == 12), 1 if start.month == 12 else start.month + 1, 1)
rows = db.session.execute(
db.select(Transaction.kind, func.sum(Transaction.amount))
.where(Transaction.user_id == current_user.id,
Transaction.happened_on >= start, Transaction.happened_on < end)
.group_by(Transaction.kind)
).all()
totals = {"income": Decimal("0"), "expense": Decimal("0")}
totals.update({kind: Decimal(total) for kind, total in rows})
return {"month": month, "income": str(totals["income"]),
"expense": str(totals["expense"]),
"balance": str(totals["income"] totals["expense"])}

赞(0)
未经允许不得转载:171主机测评 » 零基础入门python24:分类 CRUD 与用户数据隔离
分享到: 更多 (0)

评论 抢沙发

  • 昵称 (必填)
  • 邮箱 (必填)
  • 网址