diff --git a/scrapy/utils/iterators.py b/scrapy/utils/iterators.py
index 789da1392..3b504e56a 100644
--- a/scrapy/utils/iterators.py
+++ b/scrapy/utils/iterators.py
@@ -22,25 +22,41 @@ def xmliter(obj, nodename):
"""
nodename_patt = re.escape(nodename)
- HEADER_START_RE = re.compile(fr'^(.*?)<\s*{nodename_patt}(?:\s|>)', re.S)
+ DOCUMENT_HEADER_RE = re.compile(r'<\?xml[^>]+>\s*', re.S)
HEADER_END_RE = re.compile(fr'<\s*/{nodename_patt}\s*>', re.S)
+ END_TAG_RE = re.compile(r'<\s*/([^\s>]+)\s*>', re.S)
+ NAMESPACE_RE = re.compile(r'((xmlns[:A-Za-z]*)=[^>\s]+)', re.S)
text = _body_or_str(obj)
- header_start = re.search(HEADER_START_RE, text)
- header_start = header_start.group(1).strip() if header_start else ''
- header_end = re_rsearch(HEADER_END_RE, text)
- header_end = text[header_end[1]:].strip() if header_end else ''
+ document_header = re.search(DOCUMENT_HEADER_RE, text)
+ document_header = document_header.group().strip() if document_header else ''
+ header_end_idx = re_rsearch(HEADER_END_RE, text)
+ header_end = text[header_end_idx[1]:].strip() if header_end_idx else ''
+ namespaces = {}
+ if header_end:
+ for tagname in reversed(re.findall(END_TAG_RE, header_end)):
+ tag = re.search(fr'<\s*{tagname}.*?xmlns[:=][^>]*>', text[:header_end_idx[1]], re.S)
+ if tag:
+ namespaces.update(reversed(x) for x in re.findall(NAMESPACE_RE, tag.group()))
r = re.compile(fr'<{nodename_patt}[\s>].*?{nodename_patt}>', re.DOTALL)
for match in r.finditer(text):
- nodetext = header_start + match.group() + header_end
- yield Selector(text=nodetext, type='xml').xpath('//' + nodename)[0]
+ nodetext = (
+ document_header
+ + match.group().replace(
+ nodename,
+ f'{nodename} {" ".join(namespaces.values())}',
+ 1
+ )
+ + header_end
+ )
+ yield Selector(text=nodetext, type='xml')
def xmliter_lxml(obj, nodename, namespace=None, prefix='x'):
from lxml import etree
reader = _StreamReader(obj)
- tag = f'{{{namespace}}}{nodename}'if namespace else nodename
+ tag = f'{{{namespace}}}{nodename}' if namespace else nodename
iterable = etree.iterparse(reader, tag=tag, encoding=reader.encoding)
selxpath = '//' + (f'{prefix}:{nodename}' if namespace else nodename)
for _, node in iterable:
diff --git a/tests/test_utils_iterators.py b/tests/test_utils_iterators.py
index 79f5a2bbe..f84cb2956 100644
--- a/tests/test_utils_iterators.py
+++ b/tests/test_utils_iterators.py
@@ -1,5 +1,6 @@
import os
+from pytest import mark
from twisted.trial import unittest
from scrapy.utils.iterators import csviter, xmliter, _body_or_str, xmliter_lxml
@@ -134,7 +135,6 @@ class XmliterTestCase(unittest.TestCase):
"""
response = XmlResponse(url='http://mydummycompany.com', body=body)
my_iter = self.xmliter(response, 'item')
-
node = next(my_iter)
node.register_namespace('g', 'http://base.google.com/ns/1.0')
self.assertEqual(node.xpath('title/text()').getall(), ['Item 1'])
@@ -150,6 +150,55 @@ class XmliterTestCase(unittest.TestCase):
self.assertEqual(node.xpath('id/text()').getall(), [])
self.assertEqual(node.xpath('price/text()').getall(), [])
+ def test_xmliter_namespaced_nodename(self):
+ body = b"""
+
+
+
+ My Dummy Company
+ http://www.mydummycompany.com
+ This is a dummy company. We do nothing.
+ -
+ Item 1
+ This is item 1
+ http://www.mydummycompany.com/items/1
+ http://www.mydummycompany.com/images/item1.jpg
+ ITEM_1
+ 400
+
+
+
+ """
+ response = XmlResponse(url='http://mydummycompany.com', body=body)
+ my_iter = self.xmliter(response, 'g:image_link')
+ node = next(my_iter)
+ node.register_namespace('g', 'http://base.google.com/ns/1.0')
+ self.assertEqual(node.xpath('text()').extract(), ['http://www.mydummycompany.com/images/item1.jpg'])
+
+ def test_xmliter_namespaced_nodename_missing(self):
+ body = b"""
+
+
+
+ My Dummy Company
+ http://www.mydummycompany.com
+ This is a dummy company. We do nothing.
+ -
+ Item 1
+ This is item 1
+ http://www.mydummycompany.com/items/1
+ http://www.mydummycompany.com/images/item1.jpg
+ ITEM_1
+ 400
+
+
+
+ """
+ response = XmlResponse(url='http://mydummycompany.com', body=body)
+ my_iter = self.xmliter(response, 'g:link_image')
+ with self.assertRaises(StopIteration):
+ next(my_iter)
+
def test_xmliter_exception(self):
body = (
''
@@ -183,6 +232,10 @@ class XmliterTestCase(unittest.TestCase):
class LxmlXmliterTestCase(XmliterTestCase):
xmliter = staticmethod(xmliter_lxml)
+ @mark.xfail(reason='known bug of the current implementation')
+ def test_xmliter_namespaced_nodename(self):
+ super().test_xmliter_namespaced_nodename()
+
def test_xmliter_iterate_namespace(self):
body = b"""