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

如何在Python单元测试中正确Mock psycopg2数据库插入

如何在单元测试中正确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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 17:35:15