Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 8 additions & 2 deletions providers/smtp/src/airflow/providers/smtp/hooks/smtp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
"""
Expand Down
42 changes: 42 additions & 0 deletions providers/smtp/tests/unit/smtp/hooks/test_smtp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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(
Expand Down
27 changes: 27 additions & 0 deletions providers/smtp/tests/unit/smtp/notifications/test_smtp.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@

from __future__ import annotations

import json
import smtplib
import tempfile
from dataclasses import dataclass
from unittest import mock
Expand Down Expand Up @@ -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)
Expand Down