From 8d705ec3023f43c593234958aea894383e7fd375 Mon Sep 17 00:00:00 2001 From: Pablo Hoffman Date: Mon, 3 Aug 2009 22:53:08 -0300 Subject: [PATCH] added ItemLoader class, an alternative implementation of ItemBuilder with a slightly different API --- scrapy/newitem/loader/__init__.py | 114 ++++++++++++++++++++ scrapy/tests/test_itemloader.py | 171 ++++++++++++++++++++++++++++++ scrapy/utils/datatypes.py | 53 +++++++++ 3 files changed, 338 insertions(+) create mode 100644 scrapy/newitem/loader/__init__.py create mode 100644 scrapy/tests/test_itemloader.py diff --git a/scrapy/newitem/loader/__init__.py b/scrapy/newitem/loader/__init__.py new file mode 100644 index 000000000..2a61120cd --- /dev/null +++ b/scrapy/newitem/loader/__init__.py @@ -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 diff --git a/scrapy/tests/test_itemloader.py b/scrapy/tests/test_itemloader.py new file mode 100644 index 000000000..2f07aded0 --- /dev/null +++ b/scrapy/tests/test_itemloader.py @@ -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) diff --git a/scrapy/utils/datatypes.py b/scrapy/utils/datatypes.py index 795c20a91..094be5e0c 100644 --- a/scrapy/utils/datatypes.py +++ b/scrapy/utils/datatypes.py @@ -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"""