diff --git a/scrapling/spiders/robotstxt.py b/scrapling/spiders/robotstxt.py index 4e8612f..c64c66e 100644 --- a/scrapling/spiders/robotstxt.py +++ b/scrapling/spiders/robotstxt.py @@ -1,6 +1,6 @@ -from asyncio import Event from urllib.parse import urlparse +from anyio import create_task_group from protego import Protego from scrapling.core._types import Dict, Optional, Callable, Awaitable @@ -21,56 +21,40 @@ class RobotsTxtManager: - Allow/Disallow directives (including wildcards and $ anchors) - Crawl-delay directives - Deduplicates concurrent robots.txt fetches for the same domain — if multiple - requests for the same domain arrive before the first fetch completes, they - all wait for that single fetch instead of triggering redundant requests. + robots.txt is a domain-level document and does not vary by session, so the + cache is keyed by domain only. The ``sid`` parameter on public methods + controls which session is used for the initial fetch if the domain is not + yet cached, but all sessions share the same parsed result afterwards. """ def __init__(self, fetch_fn: Callable[[str, str], Awaitable]): self._fetch_fn = fetch_fn - self._cache: Dict[tuple[str, str], Protego] = {} - self._inflight: Dict[tuple[str, str], Event] = {} + self._cache: Dict[str, Protego] = {} async def _get_parser(self, url: str, sid: str) -> Protego: parsed = urlparse(url) domain = parsed.netloc + + if domain in self._cache: + return self._cache[domain] + scheme = parsed.scheme or "https" - cache_key = (domain, sid) - - # Return cached parser if available - if cache_key in self._cache: - return self._cache[cache_key] - - # If a fetch is already in-flight for this domain, wait for it to complete - if cache_key in self._inflight: - await self._inflight[cache_key].wait() - return self._cache[cache_key] - - # Mark fetch as in-flight to deduplicate concurrent requests - event = Event() - self._inflight[cache_key] = event + robots_url = f"{scheme}://{domain}/robots.txt" + content = "" + try: + response = await self._fetch_fn(robots_url, sid) + if response.status == 200: + content = response.body.decode(response.encoding, errors="replace") + except Exception as e: + log.warning(f"Failed to fetch robots.txt for {domain}: {e}") try: - robots_url = f"{scheme}://{domain}/robots.txt" - content = "" - try: - response = await self._fetch_fn(robots_url, sid) - if response.status == 200: - content = response.body.decode(response.encoding, errors="replace") - except Exception as e: - log.warning(f"Failed to fetch robots.txt for {domain}: {e}") - - try: - parser = Protego.parse(content) - except Exception as e: - log.warning(f"Failed to parse robots.txt for {domain}: {e}") - parser = Protego.parse("") - - self._cache[cache_key] = parser - finally: - event.set() - del self._inflight[cache_key] + parser = Protego.parse(content) + except Exception as e: + log.warning(f"Failed to parse robots.txt for {domain}: {e}") + parser = Protego.parse("") + self._cache[domain] = parser return parser async def can_fetch(self, url: str, sid: str) -> bool: @@ -88,7 +72,7 @@ class RobotsTxtManager: Args: url: The full URL to check - sid: Session ID for fetching robots.txt + sid: Session ID for fetching robots.txt if not yet cached Returns: True if the URL can be fetched, False otherwise @@ -104,7 +88,7 @@ class RobotsTxtManager: Args: url: Any URL on the domain to check - sid: Session ID for fetching robots.txt + sid: Session ID for fetching robots.txt if not yet cached Returns: The crawl delay in seconds, or None if not specified @@ -121,7 +105,7 @@ class RobotsTxtManager: Args: url: Any URL on the domain to check - sid: Session ID for fetching robots.txt + sid: Session ID for fetching robots.txt if not yet cached Returns: A tuple of (requests, seconds) if specified, or None if not specified @@ -137,7 +121,7 @@ class RobotsTxtManager: Args: url: Any URL on the domain to check - sid: Session ID for fetching robots.txt + sid: Session ID for fetching robots.txt if not yet cached Returns: A tuple of (crawl_delay, request_rate) where crawl_delay is in seconds @@ -151,19 +135,35 @@ class RobotsTxtManager: (rate.requests, rate.seconds) if rate is not None else None, ) - def clear_cache(self, domain: Optional[str] = None, sid: Optional[str] = None) -> None: - """Clear the robots.txt cache. + async def prefetch(self, urls: list[str], sid: str) -> None: + """Pre-warm the robots.txt cache for a list of seed URLs concurrently. + + Callers are responsible for deduplicating URLs by domain before calling + this method — passing multiple URLs for the same domain will trigger + redundant fetches since no inflight deduplication exists here. Args: - domain: If specified, only clear cache for this domain - sid: If specified, only clear cache for this session ID - If both are None, clears the entire cache + urls: Seed URLs whose domains should be pre-fetched (one per domain). + sid: Session ID to use for the robots.txt fetch requests. """ - if domain is None and sid is None: + if not urls: + return + log.debug(f"Pre-fetching robots.txt for {len(urls)} domain(s)") + async with create_task_group() as tg: + for url in urls: + tg.start_soon(self._get_parser, url, sid) + + def clear_cache(self, domain: Optional[str] = None) -> None: + """Clear the robots.txt cache. + + Note: the ``sid`` parameter was removed — the cache is now keyed by + domain only, so clearing a domain evicts all sessions at once. + + Args: + domain: If specified, only clear cache for this domain. + If None, clears the entire cache. + """ + if domain is None: self._cache.clear() else: - keys_to_remove = [ - key for key in self._cache if (domain is None or key[0] == domain) and (sid is None or key[1] == sid) - ] - for key in keys_to_remove: - del self._cache[key] + self._cache.pop(domain, None) diff --git a/tests/spiders/test_robotstxt.py b/tests/spiders/test_robotstxt.py index 5a447c0..efb8328 100644 --- a/tests/spiders/test_robotstxt.py +++ b/tests/spiders/test_robotstxt.py @@ -5,7 +5,6 @@ import asyncio import pytest from scrapling.spiders.robotstxt import RobotsTxtManager -from scrapling.core._types import List, Optional # --------------------------------------------------------------------------- @@ -28,7 +27,7 @@ def make_fetch_fn(status: int = 200, content: str = "", encoding: str = "utf-8") Attaches a `.calls` list so tests can assert how many times it was invoked and with which arguments. """ - calls: List[tuple] = [] + calls: list[tuple] = [] async def _fetch(url: str, sid: str) -> MockResponse: calls.append((url, sid)) @@ -275,11 +274,6 @@ class TestGetRequestRate: assert await mgr.get_request_rate("https://example.com/", "s1") is None -# --------------------------------------------------------------------------- -# Tests: get_sitemaps -# --------------------------------------------------------------------------- - - # --------------------------------------------------------------------------- # Tests: caching behaviour # --------------------------------------------------------------------------- @@ -308,14 +302,15 @@ class TestCachingBehaviour: assert len(fetch_fn.calls) == 1 @pytest.mark.asyncio - async def test_different_sids_use_separate_cache_entries(self): + async def test_different_sids_share_cache_entry(self): + """robots.txt is domain-level — different sessions share the same cached parser.""" fetch_fn = make_fetch_fn(content=ROBOTS_BASIC) mgr = RobotsTxtManager(fetch_fn) await mgr.can_fetch("https://example.com/", "s1") await mgr.can_fetch("https://example.com/", "s2") - assert len(fetch_fn.calls) == 2 + assert len(fetch_fn.calls) == 1 @pytest.mark.asyncio async def test_different_domains_use_separate_cache_entries(self): @@ -476,37 +471,21 @@ class TestClearCache: assert len(fetch_fn.calls) == 3 @pytest.mark.asyncio - async def test_clear_by_sid_only_invalidates_that_sid(self): + async def test_clear_by_domain_invalidates_all_sessions(self): + """Clearing a domain evicts the single shared cache entry for all sessions.""" fetch_fn = make_fetch_fn(content=ROBOTS_BASIC) mgr = RobotsTxtManager(fetch_fn) await mgr.can_fetch("https://example.com/", "s1") - await mgr.can_fetch("https://example.com/", "s2") + assert len(fetch_fn.calls) == 1 + + mgr.clear_cache(domain="example.com") + + await mgr.can_fetch("https://example.com/", "s1") # refetched — cache was cleared + await mgr.can_fetch("https://example.com/", "s2") # hits the newly warm cache, no fetch + assert len(fetch_fn.calls) == 2 - mgr.clear_cache(sid="s1") - - await mgr.can_fetch("https://example.com/", "s1") # refetched - await mgr.can_fetch("https://example.com/", "s2") # still cached - - assert len(fetch_fn.calls) == 3 - - @pytest.mark.asyncio - async def test_clear_by_domain_and_sid_targets_exact_entry(self): - fetch_fn = make_fetch_fn(content=ROBOTS_BASIC) - mgr = RobotsTxtManager(fetch_fn) - - await mgr.can_fetch("https://example.com/", "s1") - await mgr.can_fetch("https://example.com/", "s2") - assert len(fetch_fn.calls) == 2 - - mgr.clear_cache(domain="example.com", sid="s1") - - await mgr.can_fetch("https://example.com/", "s1") # refetched - await mgr.can_fetch("https://example.com/", "s2") # still cached - - assert len(fetch_fn.calls) == 3 - def test_clear_nonexistent_domain_does_not_raise(self): mgr = RobotsTxtManager(make_fetch_fn()) mgr.clear_cache(domain="nevervisited.com") # should not raise @@ -535,32 +514,30 @@ class TestClearCache: # --------------------------------------------------------------------------- -# Tests: concurrent access (double-checked locking) +# Tests: concurrent access # --------------------------------------------------------------------------- -class TestConcurrency: +class TestCacheAndConcurrency: @pytest.mark.asyncio - async def test_concurrent_calls_same_domain_same_sid_deduplicated(self): - """Multiple concurrent tasks for the same domain+sid trigger only one robots.txt fetch.""" + async def test_cached_domain_not_refetched(self): + """Once a domain is cached, subsequent calls return the cached parser without fetching.""" fetch_count = 0 - async def slow_fetch(url: str, sid: str) -> MockResponse: + async def counting_fetch(url: str, sid: str) -> MockResponse: nonlocal fetch_count fetch_count += 1 - await asyncio.sleep(0.02) # simulate network latency return MockResponse(status=200, body=ROBOTS_BASIC.encode(), encoding="utf-8") - mgr = RobotsTxtManager(slow_fetch) + mgr = RobotsTxtManager(counting_fetch) - results = await asyncio.gather(*[ - mgr.can_fetch(f"https://example.com/page{i}", "s1") - for i in range(8) - ]) + # First call fetches and caches + await mgr.can_fetch("https://example.com/page1", "s1") + # Subsequent calls hit the cache + for i in range(7): + await mgr.can_fetch(f"https://example.com/page{i + 2}", "s1") - # Concurrent calls for the same domain+sid are deduplicated to a single fetch assert fetch_count == 1 - assert all(isinstance(r, bool) for r in results) @pytest.mark.asyncio async def test_concurrent_calls_different_domains_fetch_independently(self): @@ -595,21 +572,61 @@ class TestConcurrency: assert all(r is False for r in results) @pytest.mark.asyncio - async def test_different_sids_concurrent_fetch_independently(self): + async def test_different_sids_share_cache_after_first_fetch(self): + """After the first fetch, all sessions share the cached parser regardless of sid.""" fetch_count = 0 - async def slow_fetch(url: str, sid: str) -> MockResponse: + async def counting_fetch(url: str, sid: str) -> MockResponse: nonlocal fetch_count fetch_count += 1 - await asyncio.sleep(0.01) return MockResponse(status=200, body=b"", encoding="utf-8") - mgr = RobotsTxtManager(slow_fetch) + mgr = RobotsTxtManager(counting_fetch) - await asyncio.gather( - mgr.can_fetch("https://example.com/", "s1"), - mgr.can_fetch("https://example.com/", "s2"), - mgr.can_fetch("https://example.com/", "s3"), - ) + # First call fetches and caches + await mgr.can_fetch("https://example.com/", "s1") + # s2 and s3 hit the cache — no additional fetches + await mgr.can_fetch("https://example.com/", "s2") + await mgr.can_fetch("https://example.com/", "s3") - assert fetch_count == 3 + assert fetch_count == 1 + + +# --------------------------------------------------------------------------- +# Tests: prefetch +# --------------------------------------------------------------------------- + + +class TestPrefetch: + @pytest.mark.asyncio + async def test_prefetch_fetches_all_domains(self): + fetch_fn = make_fetch_fn(content=ROBOTS_BASIC) + mgr = RobotsTxtManager(fetch_fn) + + await mgr.prefetch(["https://a.com/", "https://b.com/", "https://c.com/"], "s1") + + assert len(fetch_fn.calls) == 3 + fetched = {url for url, _ in fetch_fn.calls} + assert fetched == {"https://a.com/robots.txt", "https://b.com/robots.txt", "https://c.com/robots.txt"} + + @pytest.mark.asyncio + async def test_prefetch_warms_cache_for_subsequent_calls(self): + fetch_fn = make_fetch_fn(content=ROBOTS_BASIC) + mgr = RobotsTxtManager(fetch_fn) + + await mgr.prefetch(["https://example.com/"], "s1") + assert len(fetch_fn.calls) == 1 + + # Any subsequent call for the same domain hits the cache + await mgr.can_fetch("https://example.com/products", "s1") + await mgr.can_fetch("https://example.com/products", "s2") + assert len(fetch_fn.calls) == 1 + + @pytest.mark.asyncio + async def test_prefetch_empty_list_is_noop(self): + fetch_fn = make_fetch_fn(content=ROBOTS_BASIC) + mgr = RobotsTxtManager(fetch_fn) + + await mgr.prefetch([], "s1") + + assert len(fetch_fn.calls) == 0