mirror of https://github.com/scrapy/scrapy.git
fixed subtle bug in disk-based priority queues caused by serialization errors, and added tests
This commit is contained in:
parent
cca0b91000
commit
c382f2fc8a
|
|
@ -1,19 +1,32 @@
|
||||||
import unittest
|
from twisted.trial import unittest
|
||||||
|
|
||||||
from scrapy.utils.pqueue import PriorityQueue
|
from scrapy.utils.pqueue import PriorityQueue
|
||||||
from scrapy.utils.queue import MemoryQueue
|
from scrapy.utils.queue import MemoryQueue, DiskQueue
|
||||||
|
|
||||||
|
|
||||||
class TestMemoryQueue(MemoryQueue):
|
class TestMemoryQueue(MemoryQueue):
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self, *a, **kw):
|
||||||
super(TestMemoryQueue, self).__init__()
|
super(TestMemoryQueue, self).__init__(*a, **kw)
|
||||||
self.closed = False
|
self.closed = False
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
|
super(TestMemoryQueue, self).close()
|
||||||
self.closed = True
|
self.closed = True
|
||||||
|
|
||||||
class PriorityQueueTest(unittest.TestCase):
|
|
||||||
|
class TestDiskQueue(DiskQueue):
|
||||||
|
|
||||||
|
def __init__(self, *a, **kw):
|
||||||
|
super(TestDiskQueue, self).__init__(*a, **kw)
|
||||||
|
self.closed = False
|
||||||
|
|
||||||
|
def close(self):
|
||||||
|
super(TestDiskQueue, self).close()
|
||||||
|
self.closed = True
|
||||||
|
|
||||||
|
|
||||||
|
class MemoryPriorityQueueTest(unittest.TestCase):
|
||||||
|
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
qfactory = lambda x: TestMemoryQueue()
|
qfactory = lambda x: TestMemoryQueue()
|
||||||
|
|
@ -64,6 +77,13 @@ class PriorityQueueTest(unittest.TestCase):
|
||||||
self.assertEqual(sorted(self.q.close()), [1, 2, 3])
|
self.assertEqual(sorted(self.q.close()), [1, 2, 3])
|
||||||
assert all(q.closed for q in iqueues)
|
assert all(q.closed for q in iqueues)
|
||||||
|
|
||||||
|
def test_close_return_active(self):
|
||||||
|
self.q.push('b', 1)
|
||||||
|
self.q.push('c', 2)
|
||||||
|
self.q.push('a', 3)
|
||||||
|
self.q.pop()
|
||||||
|
self.assertEqual(sorted(self.q.close()), [2, 3])
|
||||||
|
|
||||||
def test_popped_internal_queues_closed(self):
|
def test_popped_internal_queues_closed(self):
|
||||||
self.q.push('a', 3)
|
self.q.push('a', 3)
|
||||||
self.q.push('b', 1)
|
self.q.push('b', 1)
|
||||||
|
|
@ -72,3 +92,34 @@ class PriorityQueueTest(unittest.TestCase):
|
||||||
self.assertEqual(self.q.pop(), 'b')
|
self.assertEqual(self.q.pop(), 'b')
|
||||||
self.q.close()
|
self.q.close()
|
||||||
assert p1queue.closed
|
assert p1queue.closed
|
||||||
|
|
||||||
|
|
||||||
|
class DiskPriorityQueueTest(MemoryPriorityQueueTest):
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
qfactory = lambda x: TestDiskQueue(self.mktemp())
|
||||||
|
self.q = PriorityQueue(qfactory)
|
||||||
|
|
||||||
|
def test_nonserializable_object_one(self):
|
||||||
|
self.assertRaises(TypeError, self.q.push, lambda x: x, 0)
|
||||||
|
self.assertEqual(self.q.close(), [])
|
||||||
|
|
||||||
|
def test_nonserializable_object_many_close(self):
|
||||||
|
self.q.push('a', 3)
|
||||||
|
self.q.push('b', 1)
|
||||||
|
self.assertRaises(TypeError, self.q.push, lambda x: x, 0)
|
||||||
|
self.q.push('c', 2)
|
||||||
|
self.assertEqual(self.q.pop(), 'b')
|
||||||
|
self.assertEqual(sorted(self.q.close()), [2, 3])
|
||||||
|
|
||||||
|
def test_nonserializable_object_many_pop(self):
|
||||||
|
self.q.push('a', 3)
|
||||||
|
self.q.push('b', 1)
|
||||||
|
self.assertRaises(TypeError, self.q.push, lambda x: x, 0)
|
||||||
|
self.q.push('c', 2)
|
||||||
|
self.assertEqual(self.q.pop(), 'b')
|
||||||
|
self.assertEqual(self.q.pop(), 'c')
|
||||||
|
self.assertEqual(self.q.pop(), 'a')
|
||||||
|
self.assertEqual(self.q.pop(), None)
|
||||||
|
self.assertEqual(self.q.close(), [])
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -20,17 +20,14 @@ class PriorityQueue(object):
|
||||||
self.queues = {}
|
self.queues = {}
|
||||||
self.qfactory = qfactory
|
self.qfactory = qfactory
|
||||||
for p in startprios:
|
for p in startprios:
|
||||||
q = self.qfactory(p)
|
self.queues[p] = self.qfactory(p)
|
||||||
if q:
|
|
||||||
self.queues[p] = q
|
|
||||||
self.curprio = min(startprios) if startprios else None
|
self.curprio = min(startprios) if startprios else None
|
||||||
|
|
||||||
def push(self, obj, priority=0):
|
def push(self, obj, priority=0):
|
||||||
try:
|
if priority not in self.queues:
|
||||||
q = self.queues[priority]
|
self.queues[priority] = self.qfactory(priority)
|
||||||
except KeyError:
|
q = self.queues[priority]
|
||||||
self.queues[priority] = q = self.qfactory(priority)
|
q.push(obj) # this may fail (eg. serialization error)
|
||||||
q.push(obj)
|
|
||||||
if priority < self.curprio or self.curprio is None:
|
if priority < self.curprio or self.curprio is None:
|
||||||
self.curprio = priority
|
self.curprio = priority
|
||||||
|
|
||||||
|
|
@ -39,17 +36,20 @@ class PriorityQueue(object):
|
||||||
return
|
return
|
||||||
q = self.queues[self.curprio]
|
q = self.queues[self.curprio]
|
||||||
m = q.pop()
|
m = q.pop()
|
||||||
if not q:
|
if len(q) == 0:
|
||||||
q = self.queues.pop(self.curprio)
|
del self.queues[self.curprio]
|
||||||
q.close()
|
q.close()
|
||||||
prios = self.queues.keys()
|
prios = [p for p, q in self.queues.items() if len(q) > 0]
|
||||||
self.curprio = min(prios) if prios else None
|
self.curprio = min(prios) if prios else None
|
||||||
return m
|
return m
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
for q in self.queues.values():
|
active = []
|
||||||
|
for p, q in self.queues.items():
|
||||||
|
if len(q):
|
||||||
|
active.append(p)
|
||||||
q.close()
|
q.close()
|
||||||
return self.queues.keys()
|
return active
|
||||||
|
|
||||||
def __len__(self):
|
def __len__(self):
|
||||||
return sum(len(x) for x in self.queues.values()) if self.queues else 0
|
return sum(len(x) for x in self.queues.values()) if self.queues else 0
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue