diff --git a/README.md b/README.md index 5739337..10eec04 100644 --- a/README.md +++ b/README.md @@ -154,6 +154,8 @@ The `validate_email` function also accepts the following keyword arguments ### DNS timeout and cache +The timeout is a budget shared by all MX, A, AAAA and TXT lookups for one validation. When a resolver is supplied, its current `lifetime` sets the budget for each validation. Validation does not change the settings of either a supplied resolver or dnspython's default resolver. Each lookup receives the non-negative remaining budget through dnspython's `lifetime` keyword argument, so cached answers remain eligible after the budget is exhausted. Custom `resolve` overrides must accept and honour that argument; overrides accepting only the domain and record type are incompatible. This is a soft timeout: resolver backoff and processing may take longer than the budget. + When validating many email addresses or to control the timeout (the default is 15 seconds), create a caching [dns.resolver.Resolver](https://dnspython.readthedocs.io/en/latest/resolver-class.html) to reuse in each call. The `caching_resolver` function returns one easily for you: ```python diff --git a/email_validator/deliverability.py b/email_validator/deliverability.py index 3cddc77..80c96d3 100644 --- a/email_validator/deliverability.py +++ b/email_validator/deliverability.py @@ -1,6 +1,7 @@ from typing import Any, Optional, TypedDict import ipaddress +import time from .exceptions import EmailUndeliverableError @@ -32,23 +33,27 @@ def validate_email_deliverability(domain: str, domain_i18n: str, timeout: Option # with deliverability information. # If no dns.resolver.Resolver was given, get dnspython's default resolver. - # Override the default resolver's timeout. This may affect other uses of - # dnspython in this process. if dns_resolver is None: from . import DEFAULT_TIMEOUT if timeout is None: timeout = DEFAULT_TIMEOUT dns_resolver = dns.resolver.get_default_resolver() - dns_resolver.lifetime = timeout elif timeout is not None: raise ValueError("It's not valid to pass both timeout and dns_resolver.") deliverability_info: DeliverabilityInfo = {} try: + budget = timeout if timeout is not None else dns_resolver.lifetime + deadline = time.monotonic() + budget + + def resolve(record_type: str) -> dns.resolver.Answer: + remaining = max(0.0, deadline - time.monotonic()) + return dns_resolver.resolve(domain, record_type, lifetime=remaining) + try: # Try resolving for MX records (RFC 5321 Section 5). - response = dns_resolver.resolve(domain, "MX") + response = resolve("MX") # For reporting, put them in priority order and remove the trailing dot in the qnames. mtas = sorted([(r.preference, str(r.exchange).rstrip('.')) for r in response]) @@ -84,7 +89,7 @@ def is_global_addr(address: Any) -> bool: return ipaddr.is_global try: - response = dns_resolver.resolve(domain, "A") + response = resolve("A") if not any(is_global_addr(r.address) for r in response): raise dns.resolver.NoAnswer # fall back to AAAA @@ -97,7 +102,7 @@ def is_global_addr(address: Any) -> bool: # If there was no A record, fall back to an AAAA record. # (It's unclear if SMTP servers actually do this.) try: - response = dns_resolver.resolve(domain, "AAAA") + response = resolve("AAAA") if not any(is_global_addr(r.address) for r in response): raise dns.resolver.NoAnswer @@ -118,7 +123,7 @@ def is_global_addr(address: Any) -> bool: # absence of an MX record, this is probably a good sign that the # domain is not used for email. try: - response = dns_resolver.resolve(domain, "TXT") + response = resolve("TXT") for rec in response: value = b"".join(rec.strings) if value.startswith(b"v=spf1 "): diff --git a/tests/test_deliverability.py b/tests/test_deliverability.py index e1307c2..bec002e 100644 --- a/tests/test_deliverability.py +++ b/tests/test_deliverability.py @@ -1,7 +1,11 @@ -from typing import Any +from typing import Any, Optional +from types import SimpleNamespace +import dns.exception +import dns.resolver import pytest import re +import time from email_validator import EmailUndeliverableError, \ validate_email, caching_resolver @@ -71,6 +75,140 @@ def test_timeout_and_resolver() -> None: validate_email_deliverability('timeout.com', 'timeout.com', timeout=1, dns_resolver=RESOLVER) +@pytest.fixture +def dns_clock(monkeypatch: pytest.MonkeyPatch) -> list[float]: + elapsed = [0.0] + monkeypatch.setattr(time, 'monotonic', lambda: elapsed[0]) + return elapsed + + +@pytest.mark.parametrize('supplied_resolver', [False, True]) +def test_dns_timeout_shared_across_queries(monkeypatch: pytest.MonkeyPatch, dns_clock: list[float], supplied_resolver: bool) -> None: + resolver = MockedDnsResponseData.create_resolver() + resolver.lifetime = 10 if supplied_resolver else 50 + calls: list[tuple[str, Optional[float]]] = [] + + def resolve(domain: str, record_type: str, *, lifetime: Optional[float] = None) -> Any: + calls.append((record_type, lifetime)) + budget = resolver.lifetime if lifetime is None else lifetime + if budget < 9: + dns_clock[0] += max(0, budget) + raise dns.exception.Timeout + dns_clock[0] += 9 + if record_type == 'AAAA': + return [SimpleNamespace(address='2001:4860:4860::8888')] + raise dns.resolver.NoAnswer + + monkeypatch.setattr(resolver, 'resolve', resolve) + monkeypatch.setattr(dns.resolver, 'get_default_resolver', lambda: resolver) + for _ in range(2): + if supplied_resolver: + response = validate_email_deliverability('deadline.example', 'deadline.example', dns_resolver=resolver) + else: + response = validate_email_deliverability('deadline.example', 'deadline.example', timeout=10) + assert response == {'unknown-deliverability': 'timeout'} + + assert calls == [('MX', 10), ('A', 1)] * 2 + assert dns_clock[0] == 20 + assert resolver.lifetime == (10 if supplied_resolver else 50) + + +@pytest.mark.parametrize('supplied_resolver', [False, True]) +@pytest.mark.parametrize('timeout', [0, -1, 10]) +def test_cached_answers_preserve_resolver_lifetime(monkeypatch: pytest.MonkeyPatch, supplied_resolver: bool, timeout: int) -> None: + resolver = MockedDnsResponseData.create_resolver() + resolver.lifetime = timeout if supplied_resolver else 50 + monkeypatch.setattr(dns.resolver, 'get_default_resolver', lambda: resolver) + if supplied_resolver: + response = validate_email_deliverability('gmail.com', 'gmail.com', dns_resolver=resolver) + else: + response = validate_email_deliverability('gmail.com', 'gmail.com', timeout=timeout) + + assert response == {'mx': [(5, 'gmail-smtp-in.l.google.com'), (10, 'alt1.gmail-smtp-in.l.google.com'), (20, 'alt2.gmail-smtp-in.l.google.com'), (30, 'alt3.gmail-smtp-in.l.google.com'), (40, 'alt4.gmail-smtp-in.l.google.com')], 'mx_fallback_type': None} + assert resolver.lifetime == (timeout if supplied_resolver else 50) + + +@pytest.mark.parametrize('supplied_resolver', [False, True]) +def test_resolver_override_forwards_remaining_lifetime(monkeypatch: pytest.MonkeyPatch, dns_clock: list[float], supplied_resolver: bool) -> None: + calls: list[tuple[str, Optional[float]]] = [] + + class RecordingResolver(dns.resolver.Resolver): + def resolve(self, qname: Any, rdtype: Any = 'A', *args: Any, **kwargs: Any) -> dns.resolver.Answer: + calls.append((rdtype, kwargs.get('lifetime'))) + dns_clock[0] += 4 + return super().resolve(qname, rdtype, *args, **kwargs) + + resolver = RecordingResolver(configure=False) + cache = MockedDnsResponseData.create_resolver().cache + resolver.cache = cache + resolver.lifetime = 10 if supplied_resolver else 50 + original_timeout = resolver.timeout + monkeypatch.setattr(dns.resolver, 'get_default_resolver', lambda: resolver) + monkeypatch.setattr('email_validator.DEFAULT_TIMEOUT', 10) + if supplied_resolver: + response = validate_email_deliverability('pages.github.com', 'pages.github.com', dns_resolver=resolver) + else: + response = validate_email_deliverability('pages.github.com', 'pages.github.com') + + assert response == {'mx': [(0, 'pages.github.com')], 'mx_fallback_type': 'A'} + assert calls == [('MX', 10), ('A', 6), ('TXT', 2)] + assert resolver.lifetime == (10 if supplied_resolver else 50) + assert resolver.timeout == original_timeout + assert resolver.cache is cache + + +@pytest.mark.parametrize('fallback', ['A', 'AAAA']) +@pytest.mark.parametrize('query_delay', [4, 10]) +def test_dns_budget_allows_cached_fallback_after_expiry(monkeypatch: pytest.MonkeyPatch, dns_clock: list[float], fallback: str, query_delay: int) -> None: + resolver = MockedDnsResponseData.create_resolver() + resolver.lifetime = 10 + calls: list[tuple[str, Optional[float]]] = [] + + def resolve(domain: str, record_type: str, *, lifetime: Optional[float] = None) -> Any: + calls.append((record_type, lifetime)) + dns_clock[0] += query_delay + if record_type == fallback: + return [SimpleNamespace(address='2001:4860:4860::8888')] + raise dns.resolver.NoAnswer + + monkeypatch.setattr(resolver, 'resolve', resolve) + response = validate_email_deliverability('deadline.example', 'deadline.example', dns_resolver=resolver) + assert response == {'mx': [(0, 'deadline.example')], 'mx_fallback_type': fallback} + if query_delay == 10: + if fallback == 'A': + assert calls == [('MX', 10), ('A', 0), ('TXT', 0)] + else: + assert calls == [('MX', 10), ('A', 0), ('AAAA', 0), ('TXT', 0)] + elif fallback == 'A': + assert calls == [('MX', 10), ('A', 6), ('TXT', 2)] + else: + assert calls == [('MX', 10), ('A', 6), ('AAAA', 2), ('TXT', 0)] + + +@pytest.mark.parametrize('fallback', ['A', 'AAAA']) +def test_dns_timeout_discards_partial_fallback(monkeypatch: pytest.MonkeyPatch, dns_clock: list[float], fallback: str) -> None: + resolver = MockedDnsResponseData.create_resolver() + resolver.lifetime = 2 if fallback == 'A' else 3 + calls: list[tuple[str, Optional[float]]] = [] + + def resolve(domain: str, record_type: str, *, lifetime: Optional[float] = None) -> Any: + calls.append((record_type, lifetime)) + if record_type == 'TXT': + raise dns.exception.Timeout + dns_clock[0] += 1 + if record_type == fallback: + return [SimpleNamespace(address='2001:4860:4860::8888')] + raise dns.resolver.NoAnswer + + monkeypatch.setattr(resolver, 'resolve', resolve) + response = validate_email_deliverability('deadline.example', 'deadline.example', dns_resolver=resolver) + assert response == {'unknown-deliverability': 'timeout'} + if fallback == 'A': + assert calls == [('MX', 2), ('A', 1), ('TXT', 0)] + else: + assert calls == [('MX', 3), ('A', 2), ('AAAA', 1), ('TXT', 0)] + + @pytest.mark.network def test_caching_dns_resolver() -> None: class TestCache: