added ItemLoader class, an alternative implementation of ItemBuilder with a slightly different API

This commit is contained in:
Pablo Hoffman 2009-08-03 22:53:08 -03:00
parent f05695d75e
commit 8d705ec302
3 changed files with 338 additions and 0 deletions

View File

@ -0,0 +1,114 @@
from collections import defaultdict
from types import UnboundMethodType
from scrapy.utils.misc import arg_to_iter
from scrapy.utils.python import get_func_args
from scrapy.utils.datatypes import MergeDict
from scrapy.newitem.models import Item
def tree_expander(*functions, **default_loader_args):
"""Create an ItemLoader expander from a list of functions using the tree
expansion algorithm described below.
The functions can optionally accept a ``loader_args`` argument which (if
present) will be used to pass loader arguments when the function is called.
The tree expansion algorithm consists in an ordered list of functions, each
of which receives one value and can return zero, one or more values (as a
list or iterable). If a function returns more than one value, the next
function in the pipeline will be called with each of those values,
potentially returning more values. Hence the name "tree expansion
algorithm".
"""
def wrap_y(f):
return lambda x, y, z: f(y)
def wrap_xy(f):
return lambda x, y, z: f(x, y)
def wrap_yz(f):
return lambda x, y, z: f(y, z)
wrapped_funcs = []
for func in functions:
if isinstance(func, UnboundMethodType):
func = func.im_func
if 'loader_args' in get_func_args(func):
wfunc = func
else:
wfunc = wrap_xy(func)
else:
if 'loader_args' in get_func_args(func):
wfunc = wrap_yz(func)
else:
wfunc = wrap_y(func)
wrapped_funcs.append(wfunc)
def _expander(loader, value, loader_args):
values = arg_to_iter(value)
largs = default_loader_args
if loader_args:
largs = MergeDict(loader_args, default_loader_args)
for func2 in wrapped_funcs:
next_values = []
for val in values:
val = func2(loader, val, largs)
next_values.extend(arg_to_iter(val))
values = next_values
return list(values)
return _expander
class ItemLoader(object):
item_class = Item
def __init__(self, **loader_args):
self._response = loader_args.get('response')
self._item = loader_args.get('item') or self.item_class()
self._loader_args = loader_args
self._values = defaultdict(list)
def add_value(self, field_name, value, **new_loader_args):
evalue = self._expand_value(field_name, value, new_loader_args)
self._values[field_name].extend(evalue)
def replace_value(self, field_name, value, **new_loader_args):
evalue = self._expand_value(field_name, value, new_loader_args)
self._values[field_name] = evalue
def get_item(self):
item = self._item
for field_name in self._values:
item[field_name] = self.get_value(field_name)
return item
def get_value(self, field_name):
values = self._values[field_name]
field = self._item.fields[field_name]
reducer = self.get_reducer(field_name)
# XXX: calling different methods based on reducer is ugly
if reducer:
return field.to_python(reducer(values))
else:
return field.from_unicode_list(values)
def get_reducer(self, field_name):
return getattr(self, 'reduce_%s' % field_name, None)
def get_expander(self, field_name):
return getattr(self, 'expand_%s' % field_name, self.expand)
def _expand_value(self, field_name, value, new_loader_args):
if new_loader_args:
loader_args = self._loader_args.copy()
loader_args.update(new_loader_args)
else: # shortcut for most common case
loader_args = self._loader_args
expander = self.get_expander(field_name)
return expander(value, loader_args=loader_args)
def expand(self, value, loader_args): # default expander
return value

View File

@ -0,0 +1,171 @@
import unittest
import string
from scrapy.newitem.loader import ItemLoader, tree_expander
from scrapy.newitem.builder import reducers
from scrapy.newitem import Item, fields
class BaseItem(Item):
name = fields.TextField()
class TestItem(BaseItem):
url = fields.TextField()
summary = fields.TextField()
class BaseItemLoader(ItemLoader):
item_class = TestItem
class TestItemLoader(BaseItemLoader):
expand_name = tree_expander(lambda v: v.title())
class DefaultedItemLoader(BaseItemLoader):
expand = tree_expander(lambda v: v[:-1])
class InheritDefaultedItemLoader(DefaultedItemLoader):
pass
class ListFieldTestItem(Item):
names = fields.ListField(fields.TextField())
class ListFieldItemLoader(ItemLoader):
item_class = ListFieldTestItem
expand_names = tree_expander(lambda v: v.title())
class ItemLoaderTest(unittest.TestCase):
def test_basic(self):
ib = TestItemLoader()
ib.add_value('name', u'marta')
self.assertEqual(ib.get_value('name'), u'Marta')
item = ib.get_item()
self.assertEqual(item['name'], u'Marta')
def test_multiple_functions(self):
class TestItemLoader(BaseItemLoader):
expand_name = tree_expander(lambda v: v.title(), lambda v: v[:-1])
ib = TestItemLoader()
ib.add_value('name', u'marta')
self.assertEqual(ib.get_value('name'), u'Mart')
item = ib.get_item()
self.assertEqual(item['name'], u'Mart')
def test_defaulted(self):
dib = DefaultedItemLoader()
dib.add_value('name', u'marta')
self.assertEqual(dib.get_value('name'), u'mart')
def test_inherited_default(self):
dib = InheritDefaultedItemLoader()
dib.add_value('name', u'marta')
self.assertEqual(dib.get_value('name'), u'mart')
def test_inheritance(self):
class ChildItemLoader(TestItemLoader):
expand_url = tree_expander(lambda v: v.lower())
ib = ChildItemLoader()
ib.add_value('url', u'HTTP://scrapy.ORG')
self.assertEqual(ib.get_value('url'), u'http://scrapy.org')
ib.add_value('name', u'marta')
self.assertEqual(ib.get_value('name'), u'Marta')
class ChildChildItemLoader(ChildItemLoader):
expand_url = tree_expander(lambda v: v.upper())
expand_summary = tree_expander(lambda v: v)
ib = ChildChildItemLoader()
ib.add_value('url', u'http://scrapy.org')
self.assertEqual(ib.get_value('url'), u'HTTP://SCRAPY.ORG')
ib.add_value('name', u'marta')
self.assertEqual(ib.get_value('name'), u'Marta')
def test_multiplevaluedadaptor(self):
ib = ListFieldItemLoader()
ib.add_value('names', [u'name1', u'name2'])
self.assertEqual(ib.get_value('names'), [u'Name1', u'Name2'])
def test_identity(self):
class IdentityDefaultedItemLoader(DefaultedItemLoader):
expand_name = tree_expander()
ib = IdentityDefaultedItemLoader()
ib.add_value('name', u'marta')
self.assertEqual(ib.get_value('name'), u'marta')
def test_staticmethods(self):
class ChildItemLoader(TestItemLoader):
expand_name = tree_expander(TestItemLoader.expand_name, string.swapcase)
ib = ChildItemLoader()
ib.add_value('name', u'marta')
self.assertEqual(ib.get_value('name'), u'mARTA')
def test_staticdefaults(self):
class ChildDefaultedItemLoader(DefaultedItemLoader):
expand_name = tree_expander(DefaultedItemLoader.expand, string.swapcase)
ib = ChildDefaultedItemLoader()
ib.add_value('name', u'marta')
self.assertEqual(ib.get_value('name'), u'MART')
def test_reducer(self):
ib = TestItemLoader()
ib.add_value('name', [u'mar', u'ta'])
self.assertEqual(ib.get_value('name'), u'Mar Ta')
class TakeFirstItemLoader(TestItemLoader):
reduce_name = staticmethod(reducers.take_first)
ib = TakeFirstItemLoader()
ib.add_value('name', [u'mar', u'ta'])
self.assertEqual(ib.get_value('name'), u'Mar')
def test_loader_args(self):
def expander_func_with_args(value, loader_args):
if 'val' in loader_args:
return loader_args['val']
return value
class ChildItemLoader(TestItemLoader):
expand_url = tree_expander(expander_func_with_args)
ib = ChildItemLoader(val=u'val')
ib.add_value('url', u'text')
self.assertEqual(ib.get_value('url'), 'val')
ib = ChildItemLoader()
ib.add_value('url', u'text', val=u'val')
self.assertEqual(ib.get_value('url'), 'val')
def test_add_value_unknown_field(self):
ib = TestItemLoader()
ib.add_value('wrong_field', [u'lala', u'lolo'])
self.assertRaises(KeyError, ib.get_item)

View File

@ -209,6 +209,59 @@ class CaselessDict(dict):
return dict.pop(self, self.normkey(key), *args)
class MergeDict(object):
"""
A simple class for creating new "virtual" dictionaries that actually look
up values in more than one dictionary, passed in the constructor.
If a key appears in more than one of the given dictionaries, only the
first occurrence will be used.
"""
def __init__(self, *dicts):
self.dicts = dicts
def __getitem__(self, key):
for dict_ in self.dicts:
try:
return dict_[key]
except KeyError:
pass
raise KeyError
def __copy__(self):
return self.__class__(*self.dicts)
def get(self, key, default=None):
try:
return self[key]
except KeyError:
return default
def getlist(self, key):
for dict_ in self.dicts:
if key in dict_.keys():
return dict_.getlist(key)
return []
def items(self):
item_list = []
for dict_ in self.dicts:
item_list.extend(dict_.items())
return item_list
def has_key(self, key):
for dict_ in self.dicts:
if key in dict_:
return True
return False
__contains__ = has_key
def copy(self):
"""Returns a copy of this object."""
return self.__copy__()
class PriorityQueue(object):
"""Priority queue using a deque for priority 0"""