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.queue import MemoryQueue
|
||||
from scrapy.utils.queue import MemoryQueue, DiskQueue
|
||||
|
||||
|
||||
class TestMemoryQueue(MemoryQueue):
|
||||
|
||||
def __init__(self):
|
||||
super(TestMemoryQueue, self).__init__()
|
||||
def __init__(self, *a, **kw):
|
||||
super(TestMemoryQueue, self).__init__(*a, **kw)
|
||||
self.closed = False
|
||||
|
||||
def close(self):
|
||||
super(TestMemoryQueue, self).close()
|
||||
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):
|
||||
qfactory = lambda x: TestMemoryQueue()
|
||||
|
|
@ -64,6 +77,13 @@ class PriorityQueueTest(unittest.TestCase):
|
|||
self.assertEqual(sorted(self.q.close()), [1, 2, 3])
|
||||
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):
|
||||
self.q.push('a', 3)
|
||||
self.q.push('b', 1)
|
||||
|
|
@ -72,3 +92,34 @@ class PriorityQueueTest(unittest.TestCase):
|
|||
self.assertEqual(self.q.pop(), 'b')
|
||||
self.q.close()
|
||||
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.qfactory = qfactory
|
||||
for p in startprios:
|
||||
q = self.qfactory(p)
|
||||
if q:
|
||||
self.queues[p] = q
|
||||
self.queues[p] = self.qfactory(p)
|
||||
self.curprio = min(startprios) if startprios else None
|
||||
|
||||
def push(self, obj, priority=0):
|
||||
try:
|
||||
q = self.queues[priority]
|
||||
except KeyError:
|
||||
self.queues[priority] = q = self.qfactory(priority)
|
||||
q.push(obj)
|
||||
if priority not in self.queues:
|
||||
self.queues[priority] = self.qfactory(priority)
|
||||
q = self.queues[priority]
|
||||
q.push(obj) # this may fail (eg. serialization error)
|
||||
if priority < self.curprio or self.curprio is None:
|
||||
self.curprio = priority
|
||||
|
||||
|
|
@ -39,17 +36,20 @@ class PriorityQueue(object):
|
|||
return
|
||||
q = self.queues[self.curprio]
|
||||
m = q.pop()
|
||||
if not q:
|
||||
q = self.queues.pop(self.curprio)
|
||||
if len(q) == 0:
|
||||
del self.queues[self.curprio]
|
||||
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
|
||||
return m
|
||||
|
||||
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()
|
||||
return self.queues.keys()
|
||||
return active
|
||||
|
||||
def __len__(self):
|
||||
return sum(len(x) for x in self.queues.values()) if self.queues else 0
|
||||
|
|
|
|||
Loading…
Reference in New Issue