"""Tests for db module.""" from typing import Any from unittest.mock import MagicMock, patch import pytest from pymysql.cursors import DictCursor from sync_contract_sap.db import mysql_connection @patch('sync_contract_sap.db.pymysql.connect') def test_mysql_connection_yields_connection(mock_connect: Any) -> None: """Test mysql_connection yields the pymysql connection.""" mock_conn = MagicMock() mock_connect.return_value = mock_conn with mysql_connection(host='h', user='u', password='p', database='db') as conn: assert conn is mock_conn mock_connect.assert_called_once_with( host='h', user='u', password='p', database='db', port=3306, connect_timeout=5, cursorclass=DictCursor, ) @patch('sync_contract_sap.db.pymysql.connect') def test_mysql_connection_closes_on_exit(mock_connect: Any) -> None: """Test mysql_connection closes the connection when the block exits.""" mock_conn = MagicMock() mock_connect.return_value = mock_conn with mysql_connection(host='h', user='u', password='p', database='db'): pass mock_conn.close.assert_called_once() @patch('sync_contract_sap.db.pymysql.connect') def test_mysql_connection_closes_on_exception(mock_connect: Any) -> None: """Test mysql_connection closes the connection even when an exception is raised.""" mock_conn = MagicMock() mock_connect.return_value = mock_conn with pytest.raises(RuntimeError): with mysql_connection(host='h', user='u', password='p', database='db'): raise RuntimeError('boom') mock_conn.close.assert_called_once()