Flask接口单元测试中Mock Redis/MongoDB的最优方案
解决Flask API里模块级数据库连接的Mock问题
核心解决思路
你的API在模块加载时就直接创建了redis_client和mongo_client实例,测试时不用去Mock连接方法,直接替换模块里已有的这两个客户端实例就行。用unittest.mock或者pytest的mock工具,就能让接口调用Mock后的实例,不会连真实数据库。
具体实现方法
方法1:直接Mock模块里的客户端实例(无需改API代码)
用pytest配合unittest.mock.patch就能实现,下面是完整测试代码:
from unittest.mock import Mock, patch from app import app def test_post_dostuff(): # 用patch替换app模块里的redis_client和mongo_client with patch('app.redis_client') as mock_redis, patch('app.mongo_client') as mock_mongo: # 设置Mock方法的返回值(不用管真实逻辑,只要让接口不报错) mock_redis.lpush.return_value = None mock_mongo.create_one.return_value = None # 发送POST请求,这里要传实际的请求数据 response = app.test_client().post('/dostuff', json={"name": "test"}) # 断言状态码正确 assert response.status_code == 200 # 检查redis和mongo的方法有没有被正确调用 mock_redis.lpush.assert_called_once_with({"name": "test"}) mock_mongo.create_one.assert_called_once_with({"name": "test"}) def test_get_dostuff(): # 先准备好要模拟返回的数据 mock_redis_result = [b"item1", b"item2"] mock_mongo_result = [{"id": 1, "content": "demo"}] with patch('app.redis_client') as mock_redis, patch('app.mongo_client') as mock_mongo: # 让Mock方法返回我们准备好的数据 mock_redis.lrange.return_value = mock_redis_result mock_mongo.find.return_value = mock_mongo_result # 发GET请求 response = app.test_client().get('/dostuff') # 断言状态码和返回值 assert response.status_code == 200 # 注意:Flask返回元组会自动处理成响应,这里要和接口返回格式对应 assert response.json == (mock_redis_result, mock_mongo_result) # 检查方法调用的参数是否正确(原代码里lrange第一个参数写错了,这里改成合理的列表名) mock_redis.lrange.assert_called_once_with("my_list", 0, -1) mock_mongo.find.assert_called_once_with({})
方法2:重构API代码(更优雅的长期方案)
如果后续还要加更多测试,建议改成依赖注入的方式,把客户端实例传给接口,或者存在Flask的配置里,这样测试替换更方便:
重构后的API代码:
import redis import pymongo from flask import Flask, request app = Flask(__name__) # 把初始化客户端的逻辑抽成函数 def init_db_clients(redis_url, mongo_url): redis_client = redis.from_url(redis_url) # 这里要指定具体的数据库和集合,原代码里漏了 mongo_db = pymongo.MongoClient(mongo_url).my_database return redis_client, mongo_db # 初始化全局客户端 redis_client, mongo_db = init_db_clients("redis url", "mongo url") @app.route('/dostuff', methods=['GET', 'POST']) def do_stuff_with_dbs(): post_data = request.get_json() if request.method == 'POST' else None if request.method == 'POST': redis_client.lpush("my_list", post_data) mongo_db.my_collection.insert_one(post_data) return "", 200 if request.method == 'GET': redis_list = redis_client.lrange("my_list", 0, -1) mongo_docs = list(mongo_db.my_collection.find({})) return (redis_list, mongo_docs), 200
重构后的测试代码:
from unittest.mock import Mock from app import app def test_get_dostuff(): # 创建Mock的客户端实例 mock_redis = Mock() mock_mongo = Mock() # 设置返回数据 mock_redis.lrange.return_value = [b"test_item"] mock_mongo.my_collection.find.return_value = [{"content": "test_data"}] # 直接替换app里的全局客户端 app.redis_client = mock_redis app.mongo_db = mock_mongo # 发请求断言 response = app.test_client().get('/dostuff') assert response.status_code == 200 assert response.json == ([b"test_item"], [{"content": "test_data"}])
要注意的小问题
原API代码里有几个小bug,测试前得修正:
post_data没从请求里获取,要加post_data = request.get_json()redis_client.lrange的第一个参数写的是list,应该是具体的列表名称字符串,比如"my_list"- MongoDB的
create_one应该是集合的方法,比如mongo_client.my_collection.insert_one()
内容的提问来源于stack exchange,提问作者Johnny Pongetti
相关产品推荐
相关产品推荐

