diff --git a/providers/smtp/src/airflow/providers/smtp/hooks/smtp.py b/providers/smtp/src/airflow/providers/smtp/hooks/smtp.py index 744fe705cb70e..b8fbce819041d 100644 --- a/providers/smtp/src/airflow/providers/smtp/hooks/smtp.py +++ b/providers/smtp/src/airflow/providers/smtp/hooks/smtp.py @@ -82,11 +82,17 @@ async def __aenter__(self) -> SmtpHook: return await self.aget_conn() def __exit__(self, exc_type, exc_val, exc_tb): - self._smtp_client.close() + try: + self._smtp_client.close() + finally: + self._smtp_client = None async def __aexit__(self, exc_type, exc_val, exc_tb): if self._smtp_client: - await self._smtp_client.quit() + try: + await self._smtp_client.quit() + finally: + self._smtp_client = None def _setup_oauth2(self) -> tuple[str, str]: """ diff --git a/providers/smtp/tests/unit/smtp/hooks/test_smtp.py b/providers/smtp/tests/unit/smtp/hooks/test_smtp.py index e3a3f557f70ea..cab96aeaed9c4 100644 --- a/providers/smtp/tests/unit/smtp/hooks/test_smtp.py +++ b/providers/smtp/tests/unit/smtp/hooks/test_smtp.py @@ -22,6 +22,7 @@ import smtplib import ssl import tempfile +from contextlib import nullcontext from email.mime.application import MIMEApplication from unittest import mock from unittest.mock import AsyncMock, Mock, call, patch @@ -83,6 +84,27 @@ def _create_fake_smtp(mock_smtplib, use_ssl=True): class TestSmtpHook: + @pytest.mark.parametrize("close_fails", [False, True]) + @patch("smtplib.SMTP_SSL", autospec=True) + def test_reconnect_after_context_exit(self, mock_smtp_ssl, close_fails): + clients = [mock.create_autospec(smtplib.SMTP, instance=True) for _ in range(2)] + mock_smtp_ssl.side_effect = clients + if close_fails: + clients[0].close.side_effect = OSError("Close failed") + hook = SmtpHook() + + with pytest.raises(OSError, match="Close failed") if close_fails else nullcontext(): + with hook: + hook.send_email_smtp(to=TO_EMAIL, subject=TEST_SUBJECT, html_content=TEST_BODY) + + with hook: + hook.send_email_smtp(to=TO_EMAIL, subject=TEST_SUBJECT, html_content=TEST_BODY) + + assert mock_smtp_ssl.call_count == 2 + for client in clients: + client.sendmail.assert_called_once() + client.close.assert_called_once() + @pytest.fixture(autouse=True) def setup_connections(self, create_connection_without_db): create_connection_without_db( @@ -610,6 +632,26 @@ def test_test_connection_handles_noop_responses(self, mock_smtplib, noop_respons class TestSmtpHookAsync: """Tests for async functionality in SmtpHook.""" + @pytest.mark.parametrize("quit_fails", [False, True]) + async def test_reconnect_after_context_exit(self, mock_get_connection, mocker, quit_fails): + clients = [self._create_fake_async_smtp(Mock()) for _ in range(2)] + mock_smtp = mocker.patch("airflow.providers.smtp.hooks.smtp.aiosmtplib.SMTP", side_effect=clients) + if quit_fails: + clients[0].quit.side_effect = OSError("Quit failed") + hook = SmtpHook() + + with pytest.raises(OSError, match="Quit failed") if quit_fails else nullcontext(): + async with hook: + await hook.asend_email_smtp(to=TO_EMAIL, subject=TEST_SUBJECT, html_content=TEST_BODY) + + async with hook: + await hook.asend_email_smtp(to=TO_EMAIL, subject=TEST_SUBJECT, html_content=TEST_BODY) + + assert mock_smtp.call_count == 2 + for client in clients: + client.sendmail.assert_awaited_once() + client.quit.assert_awaited_once() + @pytest.fixture(autouse=True) def setup_connections(self, create_connection_without_db): create_connection_without_db( diff --git a/providers/smtp/tests/unit/smtp/notifications/test_smtp.py b/providers/smtp/tests/unit/smtp/notifications/test_smtp.py index 4a5904405f146..02e8734539c44 100644 --- a/providers/smtp/tests/unit/smtp/notifications/test_smtp.py +++ b/providers/smtp/tests/unit/smtp/notifications/test_smtp.py @@ -17,6 +17,8 @@ from __future__ import annotations +import json +import smtplib import tempfile from dataclasses import dataclass from unittest import mock @@ -94,6 +96,31 @@ def rendered(self) -> str: class TestSmtpNotifier: + @mock.patch("smtplib.SMTP_SSL", autospec=True) + def test_repeated_notifications_reconnect(self, mock_smtp_ssl, monkeypatch): + monkeypatch.setenv( + f"AIRFLOW_CONN_{SMTP_CONN_ID.upper()}", + json.dumps( + { + "conn_type": "smtp", + "host": "smtp.example.com", + "port": 465, + "extra": {"disable_tls": True}, + } + ), + ) + clients = [mock.create_autospec(smtplib.SMTP, instance=True) for _ in range(2)] + mock_smtp_ssl.side_effect = clients + notifier = SmtpNotifier(**NOTIFIER_DEFAULT_PARAMS) + + notifier.notify({}) + notifier.notify({}) + + assert mock_smtp_ssl.call_count == 2 + for client in clients: + client.sendmail.assert_called_once() + client.close.assert_called_once() + @mock.patch("airflow.providers.smtp.notifications.smtp.SmtpHook") def test_notifier(_self, mock_smtphook_hook, create_dag_without_db): notifier = send_smtp_notification(**NOTIFIER_DEFAULT_PARAMS)