mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): run SMTP send_email off the event loop with a connection timeout (#38473)
* fix(proxy): run SMTP send_email off the event loop with a connection timeout Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): format utils.py and update _create_smtp_connection tests for timeout Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep malformed SMTP_TIMEOUT inside the email error boundary Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: retrigger ci Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci: exclude misaligned circleci coverage flag from merged codecov report Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: retrigger ci for codecov and benchmarks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci: disable carryforward for the circleci codecov flag Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci: exclude carried-forward coverage from the codecov patch status Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci: stop carrying forward the dead circleci codecov flag Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
b3235fa786
commit
3e2999f29f
5 changed files with 103 additions and 34 deletions
|
|
@ -25,6 +25,8 @@ flag_management:
|
|||
carryforward: false
|
||||
- name: proxy-db-schema-migration
|
||||
carryforward: false
|
||||
- name: circleci
|
||||
carryforward: false
|
||||
|
||||
component_management:
|
||||
individual_components:
|
||||
|
|
|
|||
|
|
@ -6031,10 +6031,42 @@ def _should_use_smtp_ssl(smtp_port: int) -> bool:
|
|||
return os.getenv("SMTP_USE_SSL", "False") == "True" or smtp_port == 465
|
||||
|
||||
|
||||
def _create_smtp_connection(smtp_host: str, smtp_port: int) -> smtplib.SMTP:
|
||||
def _create_smtp_connection(smtp_host: str, smtp_port: int, timeout: float) -> smtplib.SMTP:
|
||||
if _should_use_smtp_ssl(smtp_port=smtp_port):
|
||||
return smtplib.SMTP_SSL(host=smtp_host, port=smtp_port, context=ssl.create_default_context())
|
||||
return smtplib.SMTP(host=smtp_host, port=smtp_port)
|
||||
return smtplib.SMTP_SSL(host=smtp_host, port=smtp_port, context=ssl.create_default_context(), timeout=timeout)
|
||||
return smtplib.SMTP(host=smtp_host, port=smtp_port, timeout=timeout)
|
||||
|
||||
|
||||
def _send_smtp_message(
|
||||
email_message: MIMEMultipart,
|
||||
smtp_host: str,
|
||||
smtp_port: int,
|
||||
smtp_username: str | None,
|
||||
smtp_password: str | None,
|
||||
sender_email: str,
|
||||
receiver_email: str,
|
||||
timeout: float,
|
||||
) -> None:
|
||||
using_ssl: Final = _should_use_smtp_ssl(smtp_port=smtp_port)
|
||||
with _create_smtp_connection(
|
||||
smtp_host=smtp_host,
|
||||
smtp_port=smtp_port,
|
||||
timeout=timeout,
|
||||
) as server:
|
||||
if not using_ssl and os.getenv("SMTP_TLS", "True") != "False":
|
||||
server.starttls(context=ssl.create_default_context())
|
||||
|
||||
if smtp_username and smtp_password:
|
||||
server.login(
|
||||
user=smtp_username,
|
||||
password=smtp_password,
|
||||
)
|
||||
|
||||
server.send_message(
|
||||
msg=email_message,
|
||||
from_addr=sender_email,
|
||||
to_addrs=receiver_email,
|
||||
)
|
||||
|
||||
|
||||
async def send_email(
|
||||
|
|
@ -6080,27 +6112,18 @@ async def send_email(
|
|||
email_message.attach(MIMEText(html, "html"))
|
||||
|
||||
try:
|
||||
using_ssl: Final = _should_use_smtp_ssl(smtp_port=smtp_port)
|
||||
with _create_smtp_connection(
|
||||
smtp_timeout: Final = float(os.getenv("SMTP_TIMEOUT", "30"))
|
||||
await asyncio.to_thread(
|
||||
_send_smtp_message,
|
||||
email_message=email_message,
|
||||
smtp_host=smtp_host,
|
||||
smtp_port=smtp_port,
|
||||
) as server:
|
||||
if not using_ssl and os.getenv("SMTP_TLS", "True") != "False":
|
||||
server.starttls(context=ssl.create_default_context())
|
||||
|
||||
# Login to your email account only if smtp_username and smtp_password are provided
|
||||
if smtp_username and smtp_password:
|
||||
server.login(
|
||||
user=smtp_username,
|
||||
password=smtp_password,
|
||||
)
|
||||
|
||||
# Send the email
|
||||
server.send_message(
|
||||
msg=email_message,
|
||||
from_addr=sender_email,
|
||||
to_addrs=receiver_email,
|
||||
)
|
||||
smtp_username=smtp_username,
|
||||
smtp_password=smtp_password,
|
||||
sender_email=sender_email,
|
||||
receiver_email=receiver_email,
|
||||
timeout=smtp_timeout,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("An error occurred while sending the email:" + str(e))
|
||||
|
|
|
|||
|
|
@ -1266,13 +1266,14 @@ class TestCreateSmtpConnection:
|
|||
patch("smtplib.SMTP_SSL") as mock_smtp_ssl,
|
||||
patch("smtplib.SMTP") as mock_smtp,
|
||||
):
|
||||
result = _create_smtp_connection(smtp_host="mail.example.com", smtp_port=465)
|
||||
result = _create_smtp_connection(smtp_host="mail.example.com", smtp_port=465, timeout=30.0)
|
||||
|
||||
mock_smtp.assert_not_called()
|
||||
assert result is mock_smtp_ssl.return_value
|
||||
_, kwargs = mock_smtp_ssl.call_args
|
||||
assert kwargs["host"] == "mail.example.com"
|
||||
assert kwargs["port"] == 465
|
||||
assert kwargs["timeout"] == 30.0
|
||||
context = kwargs["context"]
|
||||
assert isinstance(context, ssl.SSLContext)
|
||||
assert context.verify_mode == ssl.CERT_REQUIRED
|
||||
|
|
@ -1286,11 +1287,11 @@ class TestCreateSmtpConnection:
|
|||
patch("smtplib.SMTP_SSL") as mock_smtp_ssl,
|
||||
patch("smtplib.SMTP") as mock_smtp,
|
||||
):
|
||||
result = _create_smtp_connection(smtp_host="mail.example.com", smtp_port=587)
|
||||
result = _create_smtp_connection(smtp_host="mail.example.com", smtp_port=587, timeout=30.0)
|
||||
|
||||
mock_smtp_ssl.assert_not_called()
|
||||
assert result is mock_smtp.return_value
|
||||
mock_smtp.assert_called_once_with(host="mail.example.com", port=587)
|
||||
mock_smtp.assert_called_once_with(host="mail.example.com", port=587, timeout=30.0)
|
||||
|
||||
|
||||
class TestSendEmailStartTls:
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from __future__ import annotations
|
|||
|
||||
import asyncio
|
||||
import sys
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from email.message import EmailMessage
|
||||
from pathlib import Path
|
||||
|
|
@ -320,6 +321,7 @@ class _SentMessage:
|
|||
body: Optional[str]
|
||||
starttls_called: bool
|
||||
login_args: Optional[tuple]
|
||||
thread_ident: int
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
@ -328,6 +330,7 @@ class InMemorySMTP:
|
|||
|
||||
sent: List[_SentMessage] = field(default_factory=list)
|
||||
raise_on_send: Optional[Exception] = None
|
||||
connection_kwargs: List[Dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
def server_factory(self) -> Callable[..., Any]:
|
||||
outer = self
|
||||
|
|
@ -370,10 +373,12 @@ class InMemorySMTP:
|
|||
body=body,
|
||||
starttls_called=self._starttls_called,
|
||||
login_args=self._login_args,
|
||||
thread_ident=threading.get_ident(),
|
||||
)
|
||||
)
|
||||
|
||||
def _factory(*args: Any, **kwargs: Any) -> _Conn:
|
||||
outer.connection_kwargs.append(dict(kwargs))
|
||||
return _Conn()
|
||||
|
||||
return _factory
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ Symbols pinned here:
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
|
@ -51,9 +52,7 @@ async def test_send_email_dispatches_via_smtp(in_memory_smtp: Any) -> None:
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_email_starttls_uses_ssl(
|
||||
in_memory_smtp: Any, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
async def test_send_email_starttls_uses_ssl(in_memory_smtp: Any, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("SMTP_USE_SSL", "True")
|
||||
await send_email(
|
||||
receiver_email="to@invalid",
|
||||
|
|
@ -82,9 +81,7 @@ async def test_send_email_error_missing_sender_email(
|
|||
) -> None:
|
||||
monkeypatch.delenv("SMTP_SENDER_EMAIL", raising=False)
|
||||
with pytest.raises(ValueError, match="SMTP_SENDER_EMAIL"):
|
||||
await send_email(
|
||||
receiver_email="x@y", subject="s", html="<p>h</p>"
|
||||
)
|
||||
await send_email(receiver_email="x@y", subject="s", html="<p>h</p>")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -105,6 +102,49 @@ async def test_send_email_error_missing_html() -> None:
|
|||
await send_email(receiver_email="x@y", subject="s", html=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_email_sets_connection_timeout(in_memory_smtp: Any) -> None:
|
||||
await send_email(
|
||||
receiver_email="to@invalid",
|
||||
subject="Hi",
|
||||
html="<p>x</p>",
|
||||
)
|
||||
assert in_memory_smtp.connection_kwargs[0].get("timeout") == 30.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_email_timeout_env_override(in_memory_smtp: Any, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("SMTP_TIMEOUT", "5")
|
||||
monkeypatch.setenv("SMTP_USE_SSL", "True")
|
||||
await send_email(
|
||||
receiver_email="to@invalid",
|
||||
subject="Hi",
|
||||
html="<p>x</p>",
|
||||
)
|
||||
assert in_memory_smtp.connection_kwargs[0].get("timeout") == 5.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_email_malformed_timeout_is_swallowed(in_memory_smtp: Any, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("SMTP_TIMEOUT", "30s")
|
||||
await send_email(
|
||||
receiver_email="to@invalid",
|
||||
subject="Hi",
|
||||
html="<p>x</p>",
|
||||
)
|
||||
assert in_memory_smtp.sent == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_email_runs_off_event_loop_thread(in_memory_smtp: Any) -> None:
|
||||
await send_email(
|
||||
receiver_email="to@invalid",
|
||||
subject="Hi",
|
||||
html="<p>x</p>",
|
||||
)
|
||||
assert in_memory_smtp.sent[0].thread_ident != threading.get_ident()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_email_smtp_failure_is_swallowed(
|
||||
in_memory_smtp: Any,
|
||||
|
|
@ -113,7 +153,5 @@ async def test_send_email_smtp_failure_is_swallowed(
|
|||
does not raise so a failing email never blocks the proxy.
|
||||
"""
|
||||
in_memory_smtp.raise_on_send = RuntimeError("smtp boom")
|
||||
await send_email(
|
||||
receiver_email="to@invalid", subject="Hi", html="<p>x</p>"
|
||||
)
|
||||
await send_email(receiver_email="to@invalid", subject="Hi", html="<p>x</p>")
|
||||
assert in_memory_smtp.sent == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue