一、问题背景:一个慢到被业务方投诉的接口

事情是这样的。我们有个商品详情接口 /api/v1/products/{id},FastAPI 写的,跑在 4 核 8G 的容器里,前面挂 Nginx,后面连 PostgreSQL 14 和 Redis 6.2。这个接口返回商品基本信息、SKU 列表、库存、促销标签,大概 2KB 的 JSON。

上线初期还好,QPS 也就几十。后来接入首页推荐和购物车,日请求量涨到 40 万左右,问题就来了:

  • P95 响应时间 2.3s,P99 直接破 4s
  • 数据库连接数经常打满(max_connections=100)
  • 监控上 PostgreSQL 的 CPU 长期 80%+
  • 业务方天天在群里 @ 我

我先用最笨的办法:在接口里打日志。结果发现单次请求里数据库查询次数高达 47 次——典型的 N+1。但光看日志不够,得拿数据说话,于是上了 profiling。

二、环境与版本

先把环境交代清楚,不同版本行为差异挺大:

  • Python 3.11.6
  • FastAPI 0.109.0
  • Flask 3.0.0(部分老接口还在用 Flask,后面会对比)
  • SQLAlchemy 2.0.25(2.0 的 selectinload 和 1.4 写法不同,注意)
  • asyncpg 0.29.0
  • Redis 6.2.14,redis-py 5.0.1
  • PostgreSQL 14.10
  • locust 2.20.0(压测)
  • py-spy 0.3.14,cProfile(标准库)

三、方案设计:先定位,再动手

我的思路很简单,分三步:

  1. 定位:用 py-spy 做采样,cProfile 做函数级统计,找出真正的热点
  2. 优化:针对热点分别处理——数据库查询、缓存、序列化
  3. 验证:用 locust 压测,对比优化前后的 P95/QPS

这里强调一点:不要凭感觉优化。我一开始以为是 JSON 序列化慢,差点去换 orjson,结果 profile 一跑发现序列化只占 3% 的时间。

3.1 用 py-spy 抓火焰图

py-spy 的好处是不用改代码,直接 attach 到进程:

# 安装
pip install py-spy==0.3.14

# 找到进程 PID
ps aux | grep uvicorn

# 采样 30 秒,生成火焰图
py-spy record -o profile.svg --pid 12345 --duration 30 --rate 100

火焰图一出来,问题一目了然:get_product_detail 这个函数下面,session.execute 被调用了 47 次,占了总时间的 68%。剩下 20% 在 jsonable_encoder(FastAPI 默认的序列化),12% 在 Redis 的同步调用上。

3.2 cProfile 做函数级统计

火焰图看趋势,具体数字还得靠 cProfile。我写了个脚本,直接调用接口函数:

# profile_api.py
import cProfile
import pstats
import asyncio
from app.api.products import get_product_detail

async def run():
    # 模拟真实请求,跑 100 次
    for i in range(100):
        await get_product_detail(product_id=i % 50, db=session, redis=redis_client)

if __name__ == "__main__":
    profiler = cProfile.Profile()
    profiler.enable()
    asyncio.run(run())
    profiler.disable()

    stats = pstats.Stats(profiler)
    stats.sort_stats("cumulative")
    stats.print_stats(20)  # 打印前 20 个耗时函数

输出里最扎眼的是:

ncalls  tottime  cumtime  function
   4700    1.234    8.567  sqlalchemy.orm.query.Query.all
   4700    0.987    6.543  asyncpg.protocol.execute
    100    0.123    2.345  jsonable_encoder

4700 次查询 / 100 次请求 = 47 次/请求,实锤 N+1。

四、核心实现:三处改造

4.1 数据库查询优化:干掉 N+1

原来的代码是这样的(简化版):

@app.get("/api/v1/products/{product_id}")
async def get_product_detail(product_id: int, db: AsyncSession = Depends(get_db)):
    product = await db.get(Product, product_id)
    if not product:
        raise HTTPException(404)

    # 这里开始 N+1
    skus = await db.execute(select(SKU).where(SKU.product_id == product_id))
    sku_list = skus.scalars().all()

    result = {"id": product.id, "name": product.name, "skus": []}
    for sku in sku_list:
        # 每个 SKU 查一次库存
        stock = await db.execute(select(Stock).where(Stock.sku_id == sku.id))
        # 每个 SKU 查一次促销
        promo = await db.execute(select(Promotion).where(Promotion.sku_id == sku.id))
        result["skus"].append({
            "id": sku.id,
            "price": sku.price,
            "stock": stock.scalar_one_or_none(),
            "promo": promo.scalar_one_or_none(),
        })
    return result

问题很明显:for 循环里查数据库。50 个 SKU 就是 100 次查询。

改造后用 SQLAlchemy 2.0 的 selectinload 预加载:

from sqlalchemy.orm import selectinload

@app.get("/api/v1/products/{product_id}")
async def get_product_detail(product_id: int, db: AsyncSession = Depends(get_db)):
    stmt = (
        select(Product)
        .options(
            selectinload(Product.skus).selectinload(SKU.stock),
            selectinload(Product.skus).selectinload(SKU.promotions),
        )
        .where(Product.id == product_id)
    )
    result = await db.execute(stmt)
    product = result.scalar_one_or_none()
    if not product:
        raise HTTPException(404)

    # 直接组装,不再查库
    return {
        "id": product.id,
        "name": product.name,
        "skus": [
            {
                "id": sku.id,
                "price": sku.price,
                "stock": sku.stock.quantity if sku.stock else 0,
                "promo": sku.promotions[0].name if sku.promotions else None,
            }
            for sku in product.skus
        ],
    }

selectinload 会发 3 条 SQL:1 条查商品,1 条查所有 SKU,1 条查所有库存和促销。47 次变 3 次。

踩坑:一开始用了 joinedload,结果因为 SKU 和 Promotion 是一对多,笛卡尔积把结果集撑到 3000 多行,反而更慢。一对多用 selectinload,多对一才用 joinedload,这是 SQLAlchemy 的老规矩了。

4.2 缓存策略:Redis 二级缓存 + 本地缓存

数据库优化完,P95 降到 400ms 左右。但商品详情是典型的读多写少,缓存能再砍一大刀。

我的缓存设计是三层的:

  1. 本地缓存cachetools.TTLCache,存热点商品,TTL 5 秒,抗住瞬时热点
  2. Redis 缓存:存完整 JSON,TTL 60 秒,加随机抖动防雪崩
  3. 数据库:兜底
import json
import random
from cachetools import TTLCache
from redis.asyncio import Redis

# 本地缓存,最多 1000 个商品,TTL 5 秒
local_cache = TTLCache(maxsize=1000, ttl=5)

async def get_product_cached(product_id: int, db: AsyncSession, redis: Redis):
    # 第一层:本地缓存
    if product_id in local_cache:
        return local_cache[product_id]

    # 第二层:Redis
    cache_key = f"product:detail:{product_id}"
    cached = await redis.get(cache_key)
    if cached:
        data = json.loads(cached)
        local_cache[product_id] = data
        return data

    # 第三层:数据库
    data = await get_product_detail(product_id, db)

    # 回写 Redis,TTL 60 秒 + 0~10 秒随机抖动
    ttl = 60 + random.randint(0, 10)
    await redis.setex(cache_key, ttl, json.dumps(data))
    local_cache[product_id] = data
    return data

踩坑一:一开始用 json.dumps 序列化 datetime,直接报错。后来加了个自定义 encoder:

from datetime import datetime

class DateTimeEncoder(json.JSONEncoder):
    def default(self, obj):
        if isinstance(obj, datetime):
            return obj.isoformat()
        return super().default(obj)

踩坑二:缓存穿透。有恶意请求一直查不存在的商品 ID,每次都打到数据库。加了个空值缓存:

if not product:
    await redis.setex(cache_key, 30, "null")  # 空值也缓存 30 秒
    return None

踩坑三:本地缓存和 Redis 不一致。商品更新后,本地缓存还有 5 秒的旧数据。我们的业务能接受 5 秒延迟,所以没做主动失效。如果要求强一致,可以用 Redis 的 pub/sub 广播失效消息。

4.3 连接池调参

优化完查询和缓存,数据库连接数还是紧张。默认的 SQLAlchemy 连接池是 pool_size=5, max_overflow=10,对 40 万日请求来说太小了。

调整后的配置:

from sqlalchemy.ext.asyncio import create_async_engine

engine = create_async_engine(
    "postgresql+asyncpg://user:pass@localhost/db",
    pool_size=20,           # 常驻连接
    max_overflow=30,        # 峰值可临时创建
    pool_timeout=10,        # 拿不到连接等 10 秒就报错
    pool_recycle=1800,      # 30 分钟回收,防止 PG 端断开
    pool_pre_ping=True,     # 每次取连接前 ping 一下
    echo=False,
)

pool_pre_ping=True 这个参数很关键。PostgreSQL 的 idle_in_transaction_session_timeout 默认是 0,但有些云厂商会设成 5 分钟,连接被服务端断了客户端不知道,下次用就报错。加了 pre_ping 会多一次 round trip,但稳定性提升明显。

五、压测验证:locust 实测数据

优化完不压测就是耍流氓。我用 locust 写了个脚本,模拟 500 并发用户,持续 3 分钟:

# locustfile.py
from locust import HttpUser, task, between
import random

class ProductUser(HttpUser):
    wait_time = between(0.1, 0.5)

    @task
    def get_product(self):
        # 80% 请求集中在 20% 的热点商品上,模拟真实分布
        if random.random() < 0.8:
            product_id = random.randint(1, 10)
        else:
            product_id = random.randint(1, 1000)
        self.client.get(f"/api/v1/products/{product_id}")

启动命令:

locust -f locustfile.py --host=http://localhost:8000 --users=500 --spawn-rate=50 --run-time=3m --headless

优化前后数据对比:

指标 优化前 优化后 提升
P50 890ms 42ms 21x
P95 2300ms 87ms 26x
P99 4100ms 156ms 26x
QPS 120 1400 11.7x
数据库 QPS 5600 380 14.7x
DB CPU 82% 18% -
错误率 0.3% 0% -

数据库 QPS 从 5600 降到 380,主要是因为 95% 的请求被缓存拦住了。

注意:这个 QPS 是在 4 核 8G 单实例上跑的,没做多实例。生产环境我们起了 3 个实例,QPS 能到 4000+。

六、踩坑与优化:那些文档不会告诉你的事

坑一:FastAPI 的 jsonable_encoder 比 orjson 慢 5 倍

profile 显示序列化占了 20%。FastAPI 默认用 jsonable_encoder + json.dumps,我换成了 orjson

from fastapi.responses import ORJSONResponse

@app.get("/api/v1/products/{product_id}", response_class=ORJSONResponse)
async def get_product_detail(...):
    ...

orjson 序列化 2KB JSON 大概 0.05ms,json.dumps 要 0.25ms。单次看不出,QPS 上千时差距就出来了。

坑二:Flask 同步接口在压测下直接跪

我们还有几个老接口是 Flask 写的。同样的压测,Flask 的 QPS 只有 180,P95 1.2s。原因是 Flask 默认单线程处理请求,要上 gunicorn + gevent:

gunicorn -w 4 -k gevent --worker-connections 1000 -b 0.0.0.0:5000 app:app

改完 QPS 到 650,但还是打不过 FastAPI 的 1400。所以新接口一律用 FastAPI,老接口慢慢迁。

坑三:Redis 同步调用阻塞事件循环

一开始我在 async 函数里用了 redis.Redis(同步客户端),结果事件循环被阻塞,QPS 上不去。换成 redis.asyncio.Redis 后,QPS 从 800 涨到 1400。

坑四:缓存雪崩

所有商品 TTL 都是 60 秒,整点一起失效,数据库瞬间被打爆。加随机抖动后解决:

ttl = 60 + random.randint(0, 10)  # 60~70 秒

七、总结

这次优化下来,几个关键结论:

  1. profile 先行:别猜,py-spy + cProfile 组合拳,10 分钟就能定位瓶颈
  2. N+1 是性能杀手:SQLAlchemy 2.0 的 selectinload 是神器,一对多用它,多对一用 joinedload
  3. 缓存要分层:本地缓存抗热点,Redis 抗全局,空值缓存防穿透,随机 TTL 防雪崩
  4. 连接池要调pool_sizemax_overflow 根据 QPS 算,pool_pre_ping 必开
  5. 序列化别忽视:orjson 比标准库快 5 倍,FastAPI 换 ORJSONResponse 一行搞定
  6. 压测要真实:locust 模拟真实流量分布,别用均匀分布自欺欺人

最后贴一下优化后的核心代码,完整的可以看我的 GitHub(假装有链接):

```python

app/api/products.py

from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import ORJSONResponse
from sqlalchemy import select
from sqlalchemy.orm import selectinload
from sqlalchemy.ext.asyncio import AsyncSession
import json
import random
from cachetools import TTLCache

router = APIRouter()
local_cache = TTLCache(maxsize=1000, ttl=5)

@router.get("/api/v1/products/{product_id}", response_class=ORJSONResponse)
async def get_product_detail(
product_id: int,
db: AsyncSession = Depends(get_db),
redis: Redis = Depends(get_redis),
):
# 本地缓存
if product_id in local_cache:
return local_cache[product_id]

# Redis 缓存
cache_key = f"product:detail:{product_id}"
cached = await redis.get(cache_key)
if cached:
    if cached == b"null":
        raise HTTPException(404)
    data = json.loads(cached)
    local_cache[product_id] = data