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而非caloriesstore_calories_fruit函数创建CaloriesCalculator时参数名拼写错误,disamount改为amount
内容的提问来源于stack exchange,提问作者Maria Lupi
相关产品推荐
相关产品推荐

