You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

FastAPI测试为何操作生产数据库而非Mock内存数据库?

问题排查:FastAPI测试用例意外写入生产数据库

我在为FastAPI接口编写测试用例,原本计划使用SQLite内存Mock数据库,但运行pytest时发现数据被写入了生产环境的calorieshistory表。以下是我的代码,请帮忙排查问题:


test_main.py 代码

import pytest
from FruitCaloriesApp.main import CaloriesCalculator, get_db, app
from FruitCaloriesApp.database import SessionLocal
from FruitCaloriesApp.models import FruitCalories, CaloriesHistory, Base
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker

engine = create_engine("sqlite:///:memory:")  # 内存SQLite数据库
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)

@pytest.fixture(scope="function")
def db_session():
    # 在内存数据库中创建表
    Base.metadata.create_all(bind=engine)
    db = SessionLocal()

    # 插入测试数据
    apple = FruitCalories(fruit="apple", calories=20.3)
    pear= FruitCalories(fruit="pear", calories=10.3)

    db.add_all([apple, pear])
    db.commit()
    yield db  
    db.close()
    Base.metadata.drop_all(bind=engine)

def test_calc_cals_apple(db_session):
    calculator = CaloriesCalculator(amount=3, fruit="apple", db=db_session)
    expected = 3 * 20.3
    assert calculator.calc_calories() == pytest.approx(expected, 0.001)

# 接口测试部分
from fastapi.testclient import TestClient       
from fastapi import status
from datetime import datetime

client = TestClient(app)  # 修正:使用导入的正确app实例

def test_calculate_calories(db_session):
    # 覆盖get_db依赖,让接口使用测试会话
    app.dependency_overrides[get_db] = lambda: db_session
    apple_calories = db_session.query(FruitCalories).filter(FruitCalories.fruit == "apple").first().calories
    response=client.post("/calculate_calories",json={"amount": 3, "fruit": "apple"})
    assert response.status_code== status.HTTP_200_OK
    assert response.json() =={"fruit": "apple", "amount": 3, "total_calories": 3 * 20.3}
    # 测试后恢复依赖
    app.dependency_overrides.clear()

def test_store_calories(db_session):
    app.dependency_overrides[get_db] = lambda: db_session
    response = client.post("/store_calories", json={"fruit": "apple", "amount": 3}, params={"user_id": 1})
    assert response.status_code == 201
    data = response.json()["data"]
    assert data["user_id"] == 1
    assert data["fruit"] == "apple"
    assert data["amount"] == 3
    expected = 3 * 20.3
    assert data["total_calories"] == pytest.approx(expected, 0.001)
    # 验证数据仅写入测试数据库
    history_entry = db_session.query(CaloriesHistory).filter(CaloriesHistory.user_id == 1).first()
    assert history_entry is not None
    app.dependency_overrides.clear()

main.py 代码

from pydantic import BaseModel
from typing import Annotated, Optional, List
from fastapi import FastAPI, Depends, status, HTTPException, Body, Path, Query
from sqlalchemy.orm import Session
from FruitCaloriesApp import models
from FruitCaloriesApp.models import FruitCalories, CaloriesHistory  # 补充导入CaloriesHistory
from FruitCaloriesApp.database import engine, SessionLocal
from datetime import datetime

app = FastAPI()                    
models.Base.metadata.create_all(bind=engine)

def get_db():
    db=SessionLocal()
    try:
        yield db
    finally:
        db.close()

db_dependency = Annotated[Session, Depends(get_db)]

class CaloriesInput(BaseModel):
    fruit: str
    amount: float  # 修正:原字段名错误,应为amount而非calories

class CaloriesHistoryResponse(BaseModel):   
    user_id: int
    date: str
    fruit: str
    amount: float
    total_calories: float

# OOP结构
class CaloriesCalculator():
    def __init__(self, amount: float, fruit: str, db: Session):
        self.amount = amount
        self.fruit = fruit
        self.db = db
        self.calories = self.get_calories()

    def get_calories(self):
        result = (self.db.query(FruitCalories.calories).filter(FruitCalories.fruit == self.fruit).first())
        if result:
            return result[0]
        else:
            return None

    def calc_calories(self):
        if self.calories is None:
            raise ValueError("valerror")     
        else:
            return self.amount * self.calories 

    def store_calories(self, db: Session, user_id: str, date, total_calories: float):
        calories_entry = CaloriesHistory(user_id=user_id, date=date, amount=self.amount,fruit=self.fruit,total_calories=total_calories)
        db.add(calories_entry)
        db.commit()
        db.refresh(calories_entry)
        return calories_entry

@app.post("/store_calories", status_code = status.HTTP_201_CREATED)
def store_calories_fruit(data: CaloriesInput, user_id: int, db: Session = Depends(get_db)):
    try:
        calculator = CaloriesCalculator(fruit=data.fruit, amount=data.amount, db=db)  # 修正:参数名拼写错误,disamount改为amount
        total_calories = calculator.calc_calories()
        date=datetime.now().isoformat()
        calculator.store_calories(db, user_id, date, total_calories)
        
        return {
            "message": "Stored", 
            "data": {"user_id": user_id, "date":  date, "fruit": data.fruit, "amount": data.amount,"total_calories": total_calories}
            }
    except ValueError:
        raise HTTPException(status_code=404, detail="Fruit not found")

核心问题与修复说明

1. 依赖未替换导致使用生产数据库

FastAPI接口的get_db依赖默认调用生产环境的SessionLocal,测试时必须通过app.dependency_overrides将其替换为测试用的内存数据库会话,否则接口会直接操作生产库。

2. 导入路径与实例错误

测试文件中存在重复导入和错误导入(如FruitCalcApp和CaloriesCalc),导致测试客户端绑定的是未修改的生产app实例,需确保客户端使用正确的app并在测试中覆盖依赖。

3. 代码拼写错误

main.py存在两处关键错误:

  • CaloriesInput类字段名错误,应为amount而非calories
  • store_calories_fruit函数创建CaloriesCalculator时参数名拼写错误,disamount改为amount

内容的提问来源于stack exchange,提问作者Maria Lupi

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 07:59:51