diff --git a/scrapy/exporters.py b/scrapy/exporters.py index 7e1d01a0a..6f679480d 100644 --- a/scrapy/exporters.py +++ b/scrapy/exporters.py @@ -38,7 +38,7 @@ class BaseItemExporter(object): raise NotImplementedError def serialize_field(self, field, name, value): - serializer = field.get('serializer', self._to_str_if_unicode) + serializer = field.get('serializer', lambda x: x) return serializer(value) def start_exporting(self): @@ -47,9 +47,6 @@ class BaseItemExporter(object): def finish_exporting(self): pass - def _to_str_if_unicode(self, value): - return value.encode(self.encoding) if isinstance(value, unicode) else value - def _get_serialized_fields(self, item, default_value=None, include_empty=None): """Return the fields to export as an iterable of tuples (name, serialized_value) @@ -89,7 +86,7 @@ class JsonLinesItemExporter(BaseItemExporter): self.file.write(self.encoder.encode(itemdict) + '\n') -class JsonItemExporter(JsonLinesItemExporter): +class JsonItemExporter(BaseItemExporter): def __init__(self, file, **kwargs): self._configure(kwargs, dont_fail=True) @@ -170,13 +167,17 @@ class CsvItemExporter(BaseItemExporter): self._headers_not_written = True self._join_multivalued = join_multivalued + def serialize_field(self, field, name, value): + serializer = field.get('serializer', self._to_str_if_unicode) + return serializer(value) + def _to_str_if_unicode(self, value): if isinstance(value, (list, tuple)): try: value = self._join_multivalued.join(value) except TypeError: # list in value may not contain strings pass - return super(CsvItemExporter, self)._to_str_if_unicode(value) + return value.encode(self.encoding) if isinstance(value, unicode) else value def export_item(self, item): if self._headers_not_written: @@ -251,7 +252,7 @@ class PythonItemExporter(BaseItemExporter): return dict(self._serialize_dict(value)) if hasattr(value, '__iter__'): return [self._serialize_value(v) for v in value] - return self._to_str_if_unicode(value) + return value.encode(self.encoding) if isinstance(value, unicode) else value def _serialize_dict(self, value): for key, val in six.iteritems(value): diff --git a/tests/test_exporters.py b/tests/test_exporters.py index b24633959..c84fb978a 100644 --- a/tests/test_exporters.py +++ b/tests/test_exporters.py @@ -23,7 +23,7 @@ class TestItem(Item): class BaseItemExporterTest(unittest.TestCase): def setUp(self): - self.i = TestItem(name=u'John\xa3', age='22') + self.i = TestItem(name=u'John\xa3', age=u'22') self.output = BytesIO() self.ie = self._get_exporter() @@ -55,6 +55,42 @@ class BaseItemExporterTest(unittest.TestCase): self.assertItemExportWorks(dict(self.i)) def test_serialize_field(self): + res = self.ie.serialize_field(self.i.fields['name'], 'name', self.i['name']) + self.assertEqual(res, u'John\xa3') + + res = self.ie.serialize_field(self.i.fields['age'], 'age', self.i['age']) + self.assertEqual(res, u'22') + + def test_fields_to_export(self): + ie = self._get_exporter(fields_to_export=['name']) + self.assertEqual(list(ie._get_serialized_fields(self.i)), [('name', u'John\xa3')]) + + def test_field_custom_serializer(self): + def custom_serializer(value): + return str(int(value) + 2) + + class CustomFieldItem(Item): + name = Field() + age = Field(serializer=custom_serializer) + + i = CustomFieldItem(name=u'John\xa3', age=u'22') + + ie = self._get_exporter() + self.assertEqual(ie.serialize_field(i.fields['name'], 'name', i['name']), u'John\xa3') + self.assertEqual(ie.serialize_field(i.fields['age'], 'age', i['age']), '24') + + +class MidRefactoringBaseItemExporterTest(BaseItemExporterTest): + """Class introduced just to keep old behavior of BaseItemExporterTest for the + test cases that inherit from it while we make changes to exporters one by + one -- a needed refactoring trick because the test cases are quite coupled. + + When we're done with the changes, we'll have ditched this class. + """ + def test_serialize_field(self): + if self.ie.__class__ is BaseItemExporter: + return + res = self.ie.serialize_field(self.i.fields['name'], 'name', self.i['name']) self.assertEqual(res, 'John\xc2\xa3') @@ -62,6 +98,9 @@ class BaseItemExporterTest(unittest.TestCase): self.assertEqual(res, '22') def test_fields_to_export(self): + if self.ie.__class__ is BaseItemExporter: + return + ie = self._get_exporter(fields_to_export=['name']) self.assertEqual(list(ie._get_serialized_fields(self.i)), [('name', 'John\xc2\xa3')]) @@ -71,6 +110,9 @@ class BaseItemExporterTest(unittest.TestCase): self.assertEqual(name, 'John\xa3') def test_field_custom_serializer(self): + if self.ie.__class__ is BaseItemExporter: + return + def custom_serializer(value): return str(int(value) + 2) @@ -85,7 +127,7 @@ class BaseItemExporterTest(unittest.TestCase): self.assertEqual(ie.serialize_field(i.fields['age'], 'age', i['age']), '24') -class PythonItemExporterTest(BaseItemExporterTest): +class PythonItemExporterTest(MidRefactoringBaseItemExporterTest): def _get_exporter(self, **kwargs): return PythonItemExporter(**kwargs) @@ -152,7 +194,7 @@ class PickleItemExporterTest(BaseItemExporterTest): self.assertEqual(pickle.load(f), i2) -class CsvItemExporterTest(BaseItemExporterTest): +class CsvItemExporterTest(MidRefactoringBaseItemExporterTest): def _get_exporter(self, **kwargs): return CsvItemExporter(self.output, **kwargs)