Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion requirements/test/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@ name = "dependencies"
requires-python = ">=3.10"
dependencies = [
"coverage[toml]",
"pyfakefs",
"pytest",
"pytest-randomly",
]
Expand Down
1 change: 0 additions & 1 deletion requirements/test/requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,6 @@ exceptiongroup==1.3.1 ; python_version == "3.10"
iniconfig==2.3.0 ; python_version >= "3.10"
packaging==26.2 ; python_version >= "3.10"
pluggy==1.6.0 ; python_version >= "3.10"
pyfakefs==6.2.0 ; python_version >= "3.10"
pygments==2.20.0 ; python_version >= "3.10"
pytest-randomly==4.1.0 ; python_version >= "3.10"
pytest==9.1.1 ; python_version >= "3.10"
Expand Down
6 changes: 3 additions & 3 deletions src/feedparser/sgmllib/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,19 +100,19 @@ def feed(self, data: str) -> None:
"""

self.rawdata = self.rawdata + data
self.goahead(0)
self.goahead(False)

def close(self) -> None:
"""Handle the remaining data."""
self.goahead(1)
self.goahead(True)

def error(self, message: str) -> t.NoReturn:
raise SGMLParseError(message)

# Internal -- handle data as far as reasonable. May leave state
# and data to be processed by a subsequent call. If 'end' is
# true, force handling all data as if followed by EOF marker.
def goahead(self, end: int) -> None:
def goahead(self, end: bool) -> None:
rawdata = self.rawdata
i = 0
n = len(rawdata)
Expand Down
130 changes: 130 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,130 @@
import re

import pytest

import feedparser.sgmllib as sgmllib


@pytest.fixture
def check_parse_error():
def checker(source: str):
parser = EventCollector()
with pytest.raises(sgmllib.SGMLParseError):
parser.feed(source)
parser.close()

yield checker


@pytest.fixture
def event_collector():
return EventCollector()


@pytest.fixture
def cdata_event_collector():
return CDATAEventCollector()


@pytest.fixture
def html_entity_collector():
return HTMLEntityCollector()


class EventCollector(sgmllib.SGMLParser):
def check_events(self, source, expected_events):
self._consume_source(source)
self._normalize_events()
assert self.events == expected_events

def _normalize_events(self):
# Normalize the list of events so that buffer artifacts don't
# separate runs of contiguous characters.
normalized_events = []
previous_type = None
for event in self.events:
current_type = event[0]
if current_type == previous_type == "data":
normalized_events[-1] = ("data", normalized_events[-1][1] + event[1])
else:
normalized_events.append(event)
previous_type = current_type
self.events = normalized_events

def _consume_source(self, source):
for s in source:
self.feed(s)
self.close()

def __init__(self) -> None:
self.events = []
super().__init__()

# structure markup

def unknown_starttag(self, tag, attrs):
self.events.append(("starttag", tag, attrs))

def unknown_endtag(self, tag):
self.events.append(("endtag", tag))

# all other markup

def handle_comment(self, data):
self.events.append(("comment", data))

def handle_charref(self, name):
self.events.append(("charref", name))

def handle_data(self, data):
self.events.append(("data", data))

def handle_decl(self, decl):
self.events.append(("decl", decl))

def handle_entityref(self, name):
self.events.append(("entityref", name))

def handle_pi(self, data):
self.events.append(("pi", data))

def unknown_decl(self, data):
self.events.append(("unknown decl", data))


class CDATAEventCollector(EventCollector):
def start_cdata(self, attrs):
self.events.append(("starttag", "cdata", attrs))
self.setliteral()


class HTMLEntityCollector(EventCollector):

entity_or_charref = re.compile(
"(?:&([a-zA-Z][-.a-zA-Z0-9]*)|&#(x[0-9a-zA-Z]+|[0-9]+))(;?)"
)

def convert_charref(self, name):
self.events.append(("charref", "convert", name))
if name[0] != "x":
return super().convert_charref(name)

def convert_codepoint(self, codepoint):
self.events.append(("codepoint", "convert", codepoint))
super().convert_codepoint(codepoint)

def convert_entityref(self, name):
self.events.append(("entityref", "convert", name))
return super().convert_entityref(name)

# These to record that they were called, then pass the call along
# to the default implementation so that its actions can be
# recorded.

def handle_charref(self, name):
self.events.append(("charref", name))
sgmllib.SGMLParser.handle_charref(self, name)

def handle_entityref(self, name):
self.events.append(("entityref", name))
sgmllib.SGMLParser.handle_entityref(self, name)
Loading