第 13 章 · Marshmallow 与 Pydantic 序列化层
本章目标:理解为何在 ch07 手写 jsonify 之外需要 Schema 序列化层;掌握 Marshmallow 的 dump/load、嵌套、partial 与 @validates;对照学习 Pydantic v2 BaseModel;在 api-demo 中重构 products API 使用 Schema;识别常见错误并完成自检与练习。
学时建议:4~5 小时(含 1.5 小时 Marshmallow 实操 + 1 小时 Pydantic 对照)
前置:本模块 ch01~ch12;重点复习 ch05 Product 模型、ch07 product_to_dict 与 ok/fail 响应封装、ch09 统一错误处理。
13.1 为何需要 Schema 层
ch07 用 product_to_dict 手动拼装字典。接口增多后会出现:
| 痛点 | 表现 |
|---|---|
| 重复代码 | 列表、详情、创建各写一遍字段 |
| 校验分散 | if not name 散落在每个视图 |
| 类型不一致 | Decimal 有时转 float,有时转 str |
| 文档难同步 | 字段增减后文档脱节(ch14 解决) |
| 安全泄漏 | 误输出 password_hash 等敏感字段 |
Schema 层在「ORM ↔ JSON」间建立单一真相源:
Request JSON → Schema.load() → Model → Schema.dump() → Response JSON
教学项目统一使用 api-demo、api.example.com、user-demo,不引用任何企业内部仓库或私有路径。
13.2 Marshmallow 安装与 ProductSchema
cd ~/python-learn/api-demo && source .venv/bin/activate
pip install marshmallow marshmallow-sqlalchemy
api_demo/api/schemas/product.py:
from decimal import Decimal
from marshmallow import Schema, fields, validates, ValidationError, EXCLUDE
class ProductSchema(Schema):
class Meta:
unknown = EXCLUDE # 忽略未声明字段,防 Mass Assignment
id = fields.Int(dump_only=True)
name = fields.Str(required=True)
slug = fields.Str(load_default=None)
price = fields.Decimal(as_string=True, places=2)
stock = fields.Int(load_default=0, validate=lambda n: n >= 0)
is_published = fields.Bool(load_default=False)
created_at = fields.DateTime(format="iso", dump_only=True)
@validates("name")
def validate_name(self, value, **kwargs):
if not value or not value.strip():
raise ValidationError("商品名称不能为空")
if len(value) > 120:
raise ValidationError("名称最长 120 字符")
@validates("price")
def validate_price(self, value, **kwargs):
if value is not None and value < Decimal("0"):
raise ValidationError("价格不能为负数")
| 方法 | 方向 | 用途 |
|---|---|---|
schema.dump(obj) | Python → dict | 响应序列化 |
schema.load(data) | dict → Python | 请求反序列化 + 校验 |
many=True | 批量 | 列表接口 |
13.3 重构 products API
api_demo/api/products.py:
from flask import request
from marshmallow import ValidationError
from api_demo.extensions import db
from api_demo.models.product import Product
from . import api_bp
from .schemas.product import ProductSchema
from .utils import ok, fail
product_schema = ProductSchema()
products_schema = ProductSchema(many=True)
@api_bp.get("/products")
def list_products():
page = request.args.get("page", 1, type=int)
per_page = min(request.args.get("per_page", 10, type=int), 50)
q = Product.query.filter_by(is_published=True).order_by(Product.id.desc())
pagination = q.paginate(page=page, per_page=per_page, error_out=False)
return ok(
data=products_schema.dump(pagination.items),
pagination={"page": pagination.page, "per_page": pagination.per_page,
"total": pagination.total, "pages": pagination.pages},
)
@api_bp.post("/products")
def create_product():
body = request.get_json(silent=True) or {}
try:
payload = product_schema.load(body)
except ValidationError as err:
return fail(42201, "参数校验失败", status=422, errors=err.messages)
product = Product(**payload)
db.session.add(product)
db.session.commit()
return ok(data=product_schema.dump(product), message="created", status=201)
与 ch07 对比:
| 维度 | ch07 手写 | ch13 Schema |
|---|---|---|
| 出参 | product_to_dict 手动维护 | dump 自动映射 |
| 入参校验 | 视图内 if not name | load 集中校验 |
| 错误格式 | 自定义字符串 | err.messages 字段级字典 |
| 局部更新 | 手动判断 key | partial=True |
13.4 partial 与 PATCH
@api_bp.patch("/products/<int:product_id>")
def update_product(product_id):
product = Product.query.get_or_404(product_id)
body = request.get_json(silent=True) or {}
try:
payload = product_schema.load(body, partial=True)
except ValidationError as err:
return fail(42201, "参数校验失败", status=422, errors=err.messages)
for key, value in payload.items():
setattr(product, key, value)
db.session.commit()
return ok(data=product_schema.dump(product))
| 参数 | 行为 |
|---|---|
partial=False | required=True 字段必须出现 |
partial=True | 仅校验请求中出现的字段,适合 PATCH |
partial=["name", "price"] | 仅指定字段允许部分出现 |
13.5 嵌套 Schema
# api_demo/api/schemas/user.py
class UserBriefSchema(Schema):
id = fields.Int(dump_only=True)
username = fields.Str()
class ProductDetailSchema(ProductSchema):
author = fields.Nested(UserBriefSchema, dump_only=True)
category = fields.Nested("CategorySchema", dump_only=True)
class ProductCreateSchema(ProductSchema):
category_id = fields.Int(required=True)
列表用 ProductSchema(many=True),详情用 ProductDetailSchema();创建时 load 得到 category_id,视图中 product.category_id = payload.pop("category_id")。
13.6 Pydantic v2 对照
pip install "pydantic>=2.0"