diff --git a/scrapy/pqueues.py b/scrapy/pqueues.py index 622f6bbc5..0681e6729 100644 --- a/scrapy/pqueues.py +++ b/scrapy/pqueues.py @@ -5,42 +5,11 @@ from collections import namedtuple from queuelib import PriorityQueue from scrapy.utils.reqser import request_to_dict, request_from_dict -from scrapy.core.downloader import Downloader -from scrapy.http import Request -from scrapy.signals import request_reached_downloader, request_left_downloader -from scrapy.utils.httpobj import urlparse_cached logger = logging.getLogger(__name__) -SCHEDULER_SLOT_META_KEY = Downloader.DOWNLOAD_SLOT - - -def _scheduler_slot_read(request, default=None): - return request.meta.get(SCHEDULER_SLOT_META_KEY, default) - - -def _scheduler_slot_write(request, slot): - request.meta[SCHEDULER_SLOT_META_KEY] = slot - - -def _set_scheduler_slot(request): - """ - >>> request = Request('http://example.com') - >>> _set_scheduler_slot(request) - 'example.com' - >>> _scheduler_slot_read(request) - 'example.com' - """ - slot = _scheduler_slot_read(request, None) - if slot is not None: - return slot - slot = urlparse_cached(request).hostname or '' - _scheduler_slot_write(request, slot) - return slot - - def _path_safe(text): """ Return a filesystem-safe version of a string ``text`` """ pathable_slot = "".join([c if c.isalnum() or c in '-._' else '_' @@ -138,6 +107,25 @@ class ScrapyPriorityQueue(PriorityQueue): return request +class DownloaderInterface(object): + + def __init__(self, crawler): + self.downloader = crawler.engine.downloader + + def stats(self, possible_slots): + return [(self._active_downloads(slot), slot) + for slot in possible_slots] + + def get_slot_key(self, request): + return self.downloader._get_slot_key(request, None) + + def _active_downloads(self, slot): + """ Return a number of requests in a Downloader for a given slot """ + if slot not in self.downloader.slots: + return 0 + return len(self.downloader.slots[slot].active) + + class DownloaderAwarePriorityQueue(object): """ PriorityQueue which takes Downlaoder activity in account: domains (slots) with the least amount of active downloads are dequeued @@ -170,68 +158,25 @@ class DownloaderAwarePriorityQueue(object): def pqfactory(startprios=()): return ScrapyPriorityQueue(crawler, qfactory, startprios, serialize) self._slot_pqueues = _SlotPriorityQueues(pqfactory, slot_startprios) - - self._active_downloads = {slot: 0 for slot in self._slot_pqueues.pqueues} - crawler.signals.connect(self.on_response_download, - signal=request_left_downloader) - crawler.signals.connect(self.on_request_reached_downloader, - signal=request_reached_downloader) self.serialize = serialize - - # There are two PriorityQueues at the same time (memory and disk-based), - # and they both listen to Downloader signals. To filter out signals - # coming from the other queue, each queue keeps track of its own - # requests using mark / unmark / check_mark methods. - def mark(self, request): - request.meta[self._DOWNLOADER_AWARE_PQ_ID] = id(self) - - def check_mark(self, request): - return request.meta.get(self._DOWNLOADER_AWARE_PQ_ID, None) == id(self) - - def unmark(self, request): - del request.meta[self._DOWNLOADER_AWARE_PQ_ID] + self._downloader_interface = DownloaderInterface(crawler) def pop(self): - slots = [(active_downloads, slot) - for slot, active_downloads in self._active_downloads.items() - if slot in self._slot_pqueues] + stats = self._downloader_interface.stats(self._slot_pqueues.pqueues) - if not slots: + if not stats: return - slot = min(slots)[1] + slot = min(stats)[1] request = self._slot_pqueues.pop_slot(slot) - self.mark(request) return request def push(self, request, priority): - slot = _set_scheduler_slot(request) + slot = self._downloader_interface.get_slot_key(request) priority_slot = _Priority(priority=priority, slot=slot) self._slot_pqueues.push_slot(slot, request, priority_slot) - if slot not in self._active_downloads: - self._active_downloads[slot] = 0 - - def on_response_download(self, request, spider): - if not self.check_mark(request): - return - self.unmark(request) - - slot = _scheduler_slot_read(request) - if slot not in self._active_downloads or self._active_downloads[slot] <= 0: - raise ValueError('Got response for a wrong slot "%s"' % (slot, )) - self._active_downloads[slot] -= 1 - if self._active_downloads[slot] == 0 and slot not in self._slot_pqueues: - del self._active_downloads[slot] - - def on_request_reached_downloader(self, request, spider): - if not self.check_mark(request): - return - - slot = _scheduler_slot_read(request) - self._active_downloads[slot] = self._active_downloads.get(slot, 0) + 1 def close(self): - self._active_downloads.clear() active = self._slot_pqueues.close() return {slot: [p.priority for p in startprios] for slot, startprios in active.items()} diff --git a/tests/test_scheduler.py b/tests/test_scheduler.py index 1bcc1e5a8..75c0b7530 100644 --- a/tests/test_scheduler.py +++ b/tests/test_scheduler.py @@ -1,20 +1,50 @@ import shutil import tempfile import unittest +import collections from twisted.internet import defer from twisted.trial.unittest import TestCase from scrapy.crawler import Crawler +from scrapy.core.downloader import Downloader from scrapy.core.scheduler import Scheduler from scrapy.http import Request -from scrapy.pqueues import _scheduler_slot_read, _scheduler_slot_write -from scrapy.signals import request_reached_downloader, request_left_downloader from scrapy.spiders import Spider +from scrapy.utils.httpobj import urlparse_cached from scrapy.utils.test import get_crawler from tests.mockserver import MockServer +MockEngine = collections.namedtuple('MockEngine', ['downloader']) +MockSlot = collections.namedtuple('MockSlot', ['active']) + + +class MockDownloader: + def __init__(self): + self.slots = dict() + + def _set_slot_key(self, slot, request, spider): + request.meta[Downloader.DOWNLOAD_SLOT] = slot + + def _get_slot_key(self, request, spider): + if Downloader.DOWNLOAD_SLOT in request.meta: + return request.meta[Downloader.DOWNLOAD_SLOT] + + return urlparse_cached(request).hostname or '' + + def increment(self, slot_key): + slot = self.slots.setdefault(slot_key, MockSlot(active=list())) + slot.active.append(1) + + def decrement(self, slot_key): + slot = self.slots.get(slot_key) + slot.active.pop() + + def close(self): + pass + + class MockCrawler(Crawler): def __init__(self, priority_queue_cls, jobdir): @@ -27,6 +57,7 @@ class MockCrawler(Crawler): DUPEFILTER_CLASS='scrapy.dupefilters.BaseDupeFilter' ) super(MockCrawler, self).__init__(Spider, settings) + self.engine = MockEngine(downloader=MockDownloader()) class SchedulerHandler: @@ -42,6 +73,7 @@ class SchedulerHandler: def close_scheduler(self): self.scheduler.close('finished') self.mock_crawler.stop() + self.mock_crawler.engine.downloader.close() def setUp(self): self.create_scheduler() @@ -147,11 +179,11 @@ class BaseSchedulerOnDiskTester(SchedulerHandler): class TestSchedulerInMemory(BaseSchedulerInMemoryTester, unittest.TestCase): - priority_queue_cls = 'queuelib.PriorityQueue' + priority_queue_cls = 'scrapy.pqueues.ScrapyPriorityQueue' class TestSchedulerOnDisk(BaseSchedulerOnDiskTester, unittest.TestCase): - priority_queue_cls = 'queuelib.PriorityQueue' + priority_queue_cls = 'scrapy.pqueues.ScrapyPriorityQueue' _SLOTS = [("http://foo.com/a", 'a'), @@ -172,7 +204,7 @@ class TestMigration(unittest.TestCase): def _migration(self, tmp_dir): prev_scheduler_handler = SchedulerHandler() - prev_scheduler_handler.priority_queue_cls = 'queuelib.PriorityQueue' + prev_scheduler_handler.priority_queue_cls = 'scrapy.pqueues.ScrapyPriorityQueue' prev_scheduler_handler.jobdir = tmp_dir prev_scheduler_handler.create_scheduler() @@ -196,30 +228,25 @@ class TestSchedulerWithDownloaderAwareInMemory(BaseSchedulerInMemoryTester, priority_queue_cls = 'scrapy.pqueues.DownloaderAwarePriorityQueue' def test_logic(self): + downloader = self.mock_crawler.engine.downloader for url, slot in _SLOTS: request = Request(url) - _scheduler_slot_write(request, slot) + downloader._set_slot_key(slot, request, None) self.scheduler.enqueue_request(request) slots = list() requests = list() while self.scheduler.has_pending_requests(): request = self.scheduler.next_request() - slots.append(_scheduler_slot_read(request)) - self.mock_crawler.signals.send_catch_log( - signal=request_reached_downloader, - request=request, - spider=self.spider - ) + slot = downloader._get_slot_key(request, None) + slots.append(slot) + downloader.increment(slot) requests.append(request) self.assertEqual(len(slots), len(_SLOTS)) for request in requests: - self.mock_crawler.signals.send_catch_log( - signal=request_left_downloader, - request=request, - spider=self.spider - ) + slot = downloader._get_slot_key(request, None) + self.mock_crawler.engine.downloader.decrement(slot) unique_slots = len(set(s for _, s in _SLOTS)) for i in range(0, len(_SLOTS), unique_slots): @@ -239,9 +266,11 @@ class TestSchedulerWithDownloaderAwareOnDisk(BaseSchedulerOnDiskTester, priority_queue_cls = 'scrapy.pqueues.DownloaderAwarePriorityQueue' def test_logic(self): + downloader = self.mock_crawler.engine.downloader + for url, slot in _SLOTS: request = Request(url) - _scheduler_slot_write(request, slot) + downloader._set_slot_key(slot, request, None) self.scheduler.enqueue_request(request) self.close_scheduler() @@ -249,27 +278,22 @@ class TestSchedulerWithDownloaderAwareOnDisk(BaseSchedulerOnDiskTester, slots = [] requests = [] + downloader = self.mock_crawler.engine.downloader while self.scheduler.has_pending_requests(): request = self.scheduler.next_request() - slots.append(_scheduler_slot_read(request)) - self.mock_crawler.signals.send_catch_log( - signal=request_reached_downloader, - request=request, - spider=self.spider - ) + slot = downloader._get_slot_key(request, None) + slots.append(slot) + downloader.increment(slot) requests.append(request) - self.assertEqual(self.scheduler.mqs._active_downloads, {}) self.assertEqual(len(slots), len(_SLOTS)) for request in requests: - self.mock_crawler.signals.send_catch_log( - signal=request_left_downloader, - request=request, - spider=self.spider - ) + slot = downloader._get_slot_key(request, None) + downloader.decrement(slot) _is_slots_unique(_SLOTS, slots) + self.assertEqual(sum(len(s.active) for s in downloader.slots.values()), 0) class StartUrlsSpider(Spider): @@ -277,6 +301,9 @@ class StartUrlsSpider(Spider): def __init__(self, start_urls): self.start_urls = start_urls + def parse(self, response): + pass + class TestIntegrationWithDownloaderAwareOnDisk(TestCase): def setUp(self):