fixed subtle bug in disk-based priority queues caused by serialization errors, and added tests

This commit is contained in:
Pablo Hoffman 2011-09-02 09:40:52 -03:00
parent cca0b91000
commit c382f2fc8a
2 changed files with 69 additions and 18 deletions

View File

@ -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(), [])

View File

@ -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