一、问题背景

事情是这样的,我们有个电商中台的商品详情接口 /api/v1/products/{id},用 FastAPI 写的,上线半年一直挺稳。直到运营搞了一次大促预热,流量翻了大概8倍,监控面板直接红了:

  • P99 响应时间:2.3s(SLA 要求 0:
    logger.warning(
    "path=%s sql_count=%d sql_time=%.3fs",
    request.url.path, stats["count"], stats["time"]
    )
    return response
跑一次请求,日志立刻说话:

path=/api/v1/products/8848 sql_count=47 sql_time=1.92s

**一个详情接口发 47 条 SQL**,这不明摆着 N+1 吗。

## 四、核心实现

### 4.1 干掉 N+1:从 47 条 SQL 到 3 条

看下原来的代码(简化版):

```python
@app.get("/api/v1/products/{product_id}")
def get_product(product_id: int, db: Session = Depends(get_db)):
    product = db.get(Product, product_id)
    return {
        "id": product.id,
        "name": product.name,
        "price": product.price,
        "skus": [
            {"id": s.id, "color": s.color, "size": s.size, "stock": s.stock}
            for s in product.skus          # 触发 N 次查询
        ],
        "tags": [t.name for t in product.tags],   # 又 N 次
        "shop": {                                  # 又 1 次
            "id": product.shop.id,
            "name": product.shop.name,
        },
    }

product.skusproduct.tagsproduct.shop 都是懒加载,每个商品有 20 多个 SKU,于是 1 + 20 + 20 + 1 ≈ 47 条。

改成 selectinload + joinedload,一次把关联数据拉回来:

from sqlalchemy.orm import selectinload, joinedload
from sqlalchemy import select

@app.get("/api/v1/products/{product_id}")
def get_product(product_id: int, db: Session = Depends(get_db)):
    stmt = (
        select(Product)
        .options(
            selectinload(Product.skus),
            selectinload(Product.tags),
            joinedload(Product.shop),
        )
        .where(Product.id == product_id)
    )
    product = db.execute(stmt).scalar_one_or_none()
    if not product:
        raise HTTPException(status_code=404, detail="product not found")
    return serialize_product(product)

为什么 skus 用 selectinload 而不是 joinedload?因为一对多 join 会产生笛卡尔积,SKU 和 tag 一起 join 的话行数会爆炸,selectinloadIN (...) 二次查询更合适。

4.2 加索引

SQL 从 47 条变 3 条后,单条还是慢。EXPLAIN ANALYZE 一看:

Seq Scan on skus  (cost=0.00..18420.00 rows=... ) (actual time=0.02..118.4 rows=...)
  Filter: (product_id = 8848)

skus.product_id 上没索引,全表扫。补上:

CREATE INDEX CONCURRENTLY idx_skus_product_id ON skus (product_id);
CREATE INDEX CONCURRENTLY idx_product_tags_product_id ON product_tags (product_id);

CONCURRENTLY 是为了不锁表,线上加索引必备。

4.3 Redis 缓存

数据库这块已经压到 ~80ms,但商品详情是典型的读多写极少,完全没必要每次都打库。上 Redis,缓存序列化后的响应体:

# app/services/product_cache.py
import json
from typing import Optional
import redis.asyncio as redis

CACHE_TTL = 300          # 5 分钟
CACHE_NULL_TTL = 30      # 空结果防穿透

class ProductCache:
    def __init__(self, client: redis.Redis):
        self.client = client

    @staticmethod
    def _key(product_id: int) -> str:
        return f"product:detail:v1:{product_id}"

    async def get(self, product_id: int) -> Optional[dict]:
        raw = await self.client.get(self._key(product_id))
        if raw is None:
            return None
        if raw == b"__NULL__":
            return {"__null__": True}
        return json.loads(raw)

    async def set(self, product_id: int, data: Optional[dict]) -> None:
        if data is None:
            await self.client.set(self._key(product_id), "__NULL__", ex=CACHE_NULL_TTL)
        else:
            await self.client.set(self._key(product_id), json.dumps(data, ensure_ascii=False), ex=CACHE_TTL)

路由里改成先查缓存:

@app.get("/api/v1/products/{product_id}")
async def get_product(product_id: int, db: Session = Depends(get_db), cache: ProductCache = Depends(get_cache)):
    cached = await cache.get(product_id)
    if cached is not None:
        if cached.get("__null__"):
            raise HTTPException(status_code=404, detail="product not found")
        return cached

    stmt = (
        select(Product)
        .options(selectinload(Product.skus), selectinload(Product.tags), joinedload(Product.shop))
        .where(Product.id == product_id)
    )
    product = db.execute(stmt).scalar_one_or_none()
    if not product:
        await cache.set(product_id, None)
        raise HTTPException(status_code=404, detail="product not found")

    data = serialize_product(product)
    await cache.set(product_id, data)
    return data

注意缓存版本号 v1:以后字段结构变了,直接升 v2,老 key 自然过期,不用手动清。

4.4 连接池调优

之前连接池是默认配置,pool_size=5, max_overflow=10,4 个 worker 一共最多 60 个连接,压测时经常 TimeoutError: QueuePool limit。改成:

# app/db.py
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker

engine = create_engine(
    "postgresql+psycopg2://user:pwd@pg:5432/shop",
    pool_size=20,
    max_overflow=10,
    pool_pre_ping=True,     # 防止连接被中间件掐断
    pool_recycle=1800,      # 30 分钟回收
    pool_timeout=5,         # 拿不到连接快速失败,别把请求堆住
    echo=False,
)
SessionLocal = sessionmaker(bind=engine, autoflush=False, autocommit=False)

五、踩坑与优化

坑1:缓存击穿。大促时有几个爆款商品的 key 同时过期,瞬间几百个请求全打到数据库。加了个简单的本地锁 + 单飞(singleflight):

import asyncio
from collections import defaultdict

_locks: dict[str, asyncio.Lock] = defaultdict(asyncio.Lock)

async def get_product_singleflight(product_id, db, cache):
    cached = await cache.get(product_id)
    if cached is not None:
        return cached
    async with _locks[f"product:{product_id}"]:
        cached = await cache.get(product_id)   # double check
        if cached is not None:
            return cached
        data = load_from_db(product_id, db)
        await cache.set(product_id, data)
        return data

坑2:selectinload + 分页。列表接口用 selectinload 时,如果不加 limit,数据库会把所有关联全查出来。列表场景应该先分页拿主表 ID,再批量查关联。

坑3:JSON 序列化成本。一开始用 Pydantic 逐字段 model_dump(),py-spy 显示序列化占了 15%。改成直接返回缓存的 dict,FastAPI 用 ORJSONResponse,序列化降到 3% 左右:

from fastapi.responses import ORJSONResponse
app = FastAPI(default_response_class=ORJSONResponse)

坑4:py-spy 采样别在生产长期开。它虽然开销小(~1%),但会 attach 到进程,建议只在排障窗口期用。

六、效果数据

locust 压测,200 并发用户,持续 5 分钟,同一台 4C8G 机器:

指标 优化前 优化后 提升
P50 1.42s 62ms 22.9x
P95 2.05s 140ms 14.6x
P99 2.30s 180ms 12.8x
QPS 50 820 16.4x
错误率 12.3% 0%
单请求 SQL 数 47 3(缓存命中时 0)
缓存命中率 94.7%
CPU 使用率 95% 41%

分阶段看:

  • 只做 N+1 优化:P99 从 2.3s → 620ms
  • 加索引:620ms → 410ms
  • 加 Redis 缓存:410ms → 180ms
  • 连接池 + ORJSON:180ms → 160ms 左右

其中 N+1 是最大头,缓存是第二大头。索引和连接池是“锦上添花”,但没有它们前面的优化效果也不稳。

七、总结

这次调优最大的感受是:别猜,去测。我一开始本能地想加机器、想换异步,结果火焰图直接告诉我瓶颈在数据库。

几点经验:

  1. py-spy 是线上排障神器,不用改代码不重启,采样完直接看火焰图。
  2. SQLAlchemy 的 before_cursor_execute 钩子是发现 N+1 最快的方法,建议在测试环境常驻。
  3. 一对多用 selectinload,多对一用 joinedload,用错了笛卡尔积会让你怀疑人生。
  4. 缓存一定要考虑穿透、击穿、雪崩,加版本号 + 空值缓存 + 单飞,基本能扛住大促。
  5. 连接池参数要跟 worker 数量匹配workers * (pool_size + max_overflow) 别超过数据库 max_connections

最后提醒一句:优化完记得再跑一次全链路压测,缓存和数据库不一致的问题往往就是在这时候暴露的。祝各位的接口都能稳在两位数毫秒。