如何在Python单元测试中正确Mock psycopg2数据库插入
我的单元测试没能Mock住psycopg2的数据库插入操作,日志和数据库都显示ID为42的记录被实际插入了,同时测试抛出断言错误:
./tests/test_requesthandler.py::TestRequestHandler::test_handle_post_coordinates Failed: [undefined]AssertionError: {'id': 42, 'latitude': 12.9, 'longitude': 77.6} != [{'latitude': 12.9, 'longitude': 77.6}]
我的单元测试代码
import unittest from unittest.mock import patch, Mock, MagicMock from http.server import BaseHTTPRequestHandler import json from src.requesthandler import RequestHandler # The code to test class TestRequestHandler(unittest.TestCase): def setUp(self): self.handler = RequestHandler() @patch('src.requesthandler.psycopg2.connect') def test_handle_post_coordinates(self, mock_connect): print(json.dumps({"latitude": "12.9", "longitude": "77.6"}).encode('utf-8')) expected = [{'latitude': 12.9, 'longitude': 77.6}] # This will disable the database connection # mock_connect.return_value.cursor.return_value.execute.return_value = None mock_con = mock_connect.return_value # result of psycopg2.connect(**connection_stuff) mock_cur = mock_con.cursor.return_value # result of con.cursor(cursor_factory=DictCursor) mock_cur.execute.return_value = expected # return this when calling cur.fetchall() mock_cur.fetchone.return_value = expected # return this when calling cur.fetchall() mock_con.commit.return_value = expected # return this when calling cur.fetchall() environ = { 'CONTENT_LENGTH': '23', 'REQUEST_METHOD': 'POST', 'PATH_INFO': '/coordinates', 'wsgi.input': Mock(read=Mock(return_value=json.dumps({'latitude': 12.9, 'longitude': 77.6}).encode('utf-8'))) } start_response = Mock() response = self.handler.handle_post_coordinates(environ, start_response) self.assertEqual(json.loads(response[0].decode().replace("'", '"')), [{'latitude': 12.9, 'longitude': 77.6}]) start_response.assert_called_with('200 OK', [('Content-type', 'text/plain')]) def test_handle_get(self): environ = { 'REQUEST_METHOD': 'GET', 'PATH_INFO': '/coordinates', } start_response = Mock() response = self.handler.handle_get(environ, start_response) self.assertEqual(json.loads(response[0].decode()), {'mssg': 'werkt123'}) start_response.assert_called_with('200 OK', [('Content-type', 'application/json')]) if __name__ == '__main__': unittest.main()
我的业务代码
import json from http.server import BaseHTTPRequestHandler, HTTPServer import psycopg2 from json import dumps from waitress import serve import logging class RequestHandler(BaseHTTPRequestHandler): # the constructor is called "__init__"for convenience def __init__(self): self.coordinates = [] print('qwe') # Connect to the PostgreSQL database self.conn = psycopg2.connect( host="localhost", database="postgres", user="postgres", password="admin" ) # Create a cursor object self.cursor = self.conn.cursor() def _send_response(self, message, status=200): self.send_response(status) self.send_header("Content-type", "application/json") self.end_headers() self.wfile.write(bytes(json.dumps(message), "utf8")) def handle_post_coordinates(self, environ, start_response): content_length = int(environ.get('CONTENT_LENGTH', 0)) request_body = environ['wsgi.input'].read(content_length).decode() coordinates = json.loads(request_body) self.coordinates.append(coordinates) self.cursor.execute("INSERT INTO coordinates (latitude, longitude) VALUES (%s, %s) RETURNING id", (coordinates['latitude'], coordinates['longitude'])) new_coordinate_id = self.cursor.fetchone()[0] self.conn.commit() new_coordinate = {'id': new_coordinate_id, 'latitude': coordinates['latitude'], 'longitude': coordinates['longitude']} status = '200 OK' headers = [('Content-type', 'text/plain')] start_response(status, headers) # return [b"Coordinates added"] return [bytes(str(new_coordinate), 'utf-8')] def handle_get(self, environ, start_response): if environ['PATH_INFO'] == '/coordinates': # self.cursor.execute("SELECT * FROM coordinates") # coordinates = self.cursor.fetchall() # coordinates = [{'id': c[0], 'latitude': c[1], 'longitude': c[2]} for c in coordinates] # response = dumps(coordinates) response = dumps({'mssg' : 'werkt123'}) # response = {'mssg' : 'haha3'} elif environ['PATH_INFO'].startswith('/coordinates/'): coordinate_id = int(environ['PATH_INFO'].split('/')[-1]) self.cursor.execute("SELECT * FROM coordinates WHERE id = %s", (coordinate_id,)) coordinate = self.cursor.fetchone() if coordinate: coordinate = {'id': coordinate[0], 'latitude': coordinate[1], 'longitude': coordinate[2]} response = dumps(coordinate) else: response = dumps({'error': 'Coordinate not found'}) start_response("404 Not Found", [('Content-type', 'application/json')]) return [response.encode()] else: response = dumps({'error': 'Invalid endpoint'}) start_response("404 Not Found", [('Content-type', 'application/json')]) return [response.encode()] start_response("200 OK", [('Content-type', 'application/json')]) return [response.encode()] def application(self, environ, start_response): try: path = environ.get('PATH_INFO', '').lstrip('/') if path == 'coordinates': if environ['REQUEST_METHOD'] == 'GET': return self.handle_get(environ, start_response) elif environ['REQUEST_METHOD'] == 'POST': return self.handle_post_coordinates(environ, start_response) # return self.do_POST(environ, start_response) # elif environ['REQUEST_METHOD'] == 'PUT': # return handle_put_coordinates(environ, start_response) # elif environ['REQUEST_METHOD'] == 'DELETE': # return handle_delete_coordinates(environ, start_response) else: start_response("404 Not Found", [('Content-type', 'text/plain')]) return [b"Not Found"] except Exception as e: print("Error:", str(e)) start_response("500 Internal Server Error", []) return [b"Error: " + str(e).encode()]
问题分析与解决办法
1. Mock时机错误:真实连接已提前建立
你的RequestHandler在__init__方法里直接创建了真实的数据库连接,而单元测试的setUp先实例化了RequestHandler,之后才触发@patch——这时候真实连接已经建立,Mock完全没生效,导致测试时实际执行了数据库插入。
解决:
要么调整Mock时机,在实例化RequestHandler前就Mockpsycopg2.connect;要么重构业务代码,支持注入数据库连接(推荐后者,更灵活)。
重构业务代码的__init__方法:
def __init__(self, conn=None): self.coordinates = [] # 测试时传入Mock连接,生产环境自动创建真实连接 if conn: self.conn = conn else: self.conn = psycopg2.connect( host="localhost", database="postgres", user="postgres", password="admin" ) self.cursor = self.conn.cursor()
2. Mock返回值配置错误
业务代码中cursor.fetchone()[0]是获取INSERT RETURNING id返回的ID值,而你给mock_cur.fetchone.return_value设置成了[{'latitude': 12.9, 'longitude': 77.6}]——真实场景下fetchone()返回的是包含ID的元组(比如(42,)),你的配置会导致new_coordinate_id拿到整个字典,和预期不符。
修正Mock配置:
# 模拟INSERT返回的ID,fetchone()返回元组(42,) mock_cur.fetchone.return_value = (42,)
3. 断言逻辑错误
业务代码返回的是单个字典{'id': 42, ...},但你的测试断言里拿它和列表[{'latitude': ...}]比较,自然会失败。
修正断言:
expected_response = {'id': 42, 'latitude': 12.9, 'longitude': 77.6} self.assertEqual(json.loads(response[0].decode().replace("'", '"')), expected_response)
4. 完整修正后的测试代码
import unittest from unittest.mock import patch, Mock, MagicMock import json from src.requesthandler import RequestHandler class TestRequestHandler(unittest.TestCase): @patch('src.requesthandler.psycopg2.connect') def test_handle_post_coordinates(self, mock_connect): # 先配置Mock,再实例化Handler mock_con = mock_connect.return_value mock_cur = mock_con.cursor.return_value # 模拟INSERT返回的ID mock_cur.fetchone.return_value = (42,) # 注入Mock连接到Handler self.handler = RequestHandler(conn=mock_con) environ = { 'CONTENT_LENGTH': '27', # 修正为json字符串的正确长度 'REQUEST_METHOD': 'POST', 'PATH_INFO': '/coordinates', 'wsgi.input': Mock(read=Mock(return_value=json.dumps({'latitude': 12.9, 'longitude': 77.6}).encode('utf-8'))) } start_response = Mock() response = self.handler.handle_post_coordinates(environ, start_response) # 修正断言预期值 expected_response = {'id': 42, 'latitude': 12.9, 'longitude': 77.6} self.assertEqual(json.loads(response[0].decode().replace("'", '"')), expected_response) start_response.assert_called_with('200 OK', [('Content-type', 'text/plain')]) # 验证执行了正确的SQL语句 mock_cur.execute.assert_called_once_with( "INSERT INTO coordinates (latitude, longitude) VALUES (%s, %s) RETURNING id", (12.9, 77.6) ) # 验证事务被提交 mock_con.commit.assert_called_once() def test_handle_get(self): # 注入Mock连接避免真实数据库操作 mock_conn = Mock() self.handler = RequestHandler(conn=mock_conn) environ = { 'REQUEST_METHOD': 'GET', 'PATH_INFO': '/coordinates', } start_response = Mock() response = self.handler.handle_get(environ, start_response) self.assertEqual(json.loads(response[0].decode()), {'mssg': 'werkt123'}) start_response.assert_called_with('200 OK', [('Content-type', 'application/json')]) if __name__ == '__main__': unittest.main()
内容的提问来源于stack exchange,提问作者superkytoz

