"""Test ssh_tunnel_forwarder module.""" from unittest.mock import patch, MagicMock import pytest from sme_labelcopy_loader import ssh_tunnel_forwarder @ssh_tunnel_forwarder.ssh_tunnel def func(data, host, port): """simple function for testing.""" return data, host, port @pytest.yield_fixture def mock_ssh(): """Mock SSHTunnelForwarder.""" path = 'sme_labelcopy_loader.ssh_tunnel_forwarder.SSHTunnelForwarder' with patch(path) as mock_sshtunnel: mock_mock_sshtunnel_context = ( mock_sshtunnel.return_value.__enter__.return_value) yield mock_mock_sshtunnel_context def test_ssh_tunnel_forwarder_with_dev_env(monkeypatch, mock_ssh): """Test ssh_tunnel_forwarder with dev environment.""" monkeypatch.setenv('Environment', 'dev') result = func('2020-10-01', '10.10.10.10', 6432) assert result == ( '2020-10-01', mock_ssh.local_bind_host, mock_ssh.local_bind_port) def test_ssh_tunnel_forwarder_with_prod_env(monkeypatch, mock_ssh): """Test ssh_tunnel_forwarder with not dev environment.""" monkeypatch.setenv('Environment', 'prod') result = func('2020-10-01', '10.10.10.10', 6432) assert result == ( '2020-10-01', '10.10.10.10', 6432)