180 lines
6 KiB
Python
180 lines
6 KiB
Python
|
|
"""The Mews subset: the allowlist and the transform enforcing it.
|
||
|
|
|
||
|
|
The allowlist is read from the DTD, the only source of what is allowed.
|
||
|
|
Element names come from <!ELEMENT> declarations, attributes from
|
||
|
|
<!ATTLIST> groups, with parameter entities expanded so groups such as
|
||
|
|
%Common.attrib; contribute the attributes they stand for.
|
||
|
|
"""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
from collections.abc import Callable
|
||
|
|
from io import BytesIO
|
||
|
|
import re
|
||
|
|
from typing import TypeAlias
|
||
|
|
|
||
|
|
from lxml import etree
|
||
|
|
|
||
|
|
XHTML_NS = "http://www.w3.org/1999/xhtml"
|
||
|
|
# lxml represents the xml:lang attribute with this Clark-notation key.
|
||
|
|
XML_NS = "{http://www.w3.org/XML/1998/namespace}"
|
||
|
|
|
||
|
|
Tree: TypeAlias = etree._ElementTree
|
||
|
|
Element: TypeAlias = etree._Element
|
||
|
|
Allowlist: TypeAlias = dict[str, set[str]]
|
||
|
|
|
||
|
|
# An attribute name followed by a type token, e.g. "href CDATA #IMPLIED".
|
||
|
|
_ATTRIBUTE = re.compile(
|
||
|
|
r"([\w.:-]+)\s+(?:CDATA|NMTOKENS?|ID|IDREFS?|ENTIT(?:IES|Y)|NOTATION|\()"
|
||
|
|
)
|
||
|
|
_ENTITY = re.compile(
|
||
|
|
r"<!ENTITY\s+%\s+([\w.-]+)\s+(?:PUBLIC|SYSTEM)?\s*\"(.*?)\"\s*>", re.DOTALL
|
||
|
|
)
|
||
|
|
_ATTLIST = re.compile(r"<!ATTLIST\s+([\w.:-]+)(.*?)>", re.DOTALL)
|
||
|
|
_ELEMENT = re.compile(r"<!ELEMENT\s+([\w.:-]+)")
|
||
|
|
_REFERENCE = re.compile(r"%([\w.-]+);")
|
||
|
|
|
||
|
|
|
||
|
|
class MewsSubsetError(RuntimeError):
|
||
|
|
"""Content outside the Mews subset, in strict mode."""
|
||
|
|
|
||
|
|
|
||
|
|
def parse_allowlist(dtd: str) -> Allowlist:
|
||
|
|
"""Map element name to its allowed attributes, read from DTD markup."""
|
||
|
|
entities = dict(_ENTITY.findall(dtd))
|
||
|
|
|
||
|
|
def expand(text: str) -> str:
|
||
|
|
for _ in range(len(entities) + 1):
|
||
|
|
expanded = _REFERENCE.sub(
|
||
|
|
lambda match: entities.get(match.group(1), match.group(0)), text
|
||
|
|
)
|
||
|
|
if expanded == text:
|
||
|
|
break
|
||
|
|
text = expanded
|
||
|
|
return text
|
||
|
|
|
||
|
|
allowlist: Allowlist = {}
|
||
|
|
for name, body in _ATTLIST.findall(dtd):
|
||
|
|
allowlist.setdefault(name, set()).update(_ATTRIBUTE.findall(expand(body)))
|
||
|
|
for name in _ELEMENT.findall(dtd):
|
||
|
|
allowlist.setdefault(name, set())
|
||
|
|
return allowlist
|
||
|
|
|
||
|
|
|
||
|
|
def element_name(tag: str) -> str:
|
||
|
|
"""Local name of a tag, with any namespace stripped."""
|
||
|
|
return tag.rpartition("}")[2] if tag.startswith("{") else tag
|
||
|
|
|
||
|
|
|
||
|
|
def attribute_name(key: str) -> str:
|
||
|
|
"""Name of an attribute key, mapping the xml namespace to the xml: prefix."""
|
||
|
|
if key.startswith(XML_NS):
|
||
|
|
return "xml:" + key[len(XML_NS) :]
|
||
|
|
return key
|
||
|
|
|
||
|
|
|
||
|
|
def parse_page(data: bytes) -> tuple[Tree, str | None]:
|
||
|
|
"""Parse one page, returning its tree and original doctype."""
|
||
|
|
parser = etree.XMLParser(resolve_entities=False, load_dtd=False)
|
||
|
|
tree = etree.parse(BytesIO(data), parser)
|
||
|
|
return tree, tree.docinfo.doctype or None
|
||
|
|
|
||
|
|
|
||
|
|
def ensure_xhtml_root(tree: Tree) -> Tree:
|
||
|
|
"""Give <html> the XHTML default namespace if it is missing.
|
||
|
|
|
||
|
|
Children keep their own tags; a default namespace declared on the
|
||
|
|
root element covers them.
|
||
|
|
"""
|
||
|
|
root = tree.getroot()
|
||
|
|
if root.nsmap.get(None) == XHTML_NS:
|
||
|
|
return tree
|
||
|
|
html = etree.Element(f"{{{XHTML_NS}}}html")
|
||
|
|
html.attrib.update(root.attrib)
|
||
|
|
html.text, html.tail = root.text, root.tail
|
||
|
|
for child in root:
|
||
|
|
html.append(child)
|
||
|
|
return etree.ElementTree(html)
|
||
|
|
|
||
|
|
|
||
|
|
def serialize(tree: Tree, doctype: str | None) -> bytes:
|
||
|
|
"""Serialize as UTF-8 XML; empty elements come out self-closed."""
|
||
|
|
return etree.tostring(tree, xml_declaration=True, encoding="UTF-8", doctype=doctype)
|
||
|
|
|
||
|
|
|
||
|
|
def enforce_subset(
|
||
|
|
tree: Tree,
|
||
|
|
allowlist: Allowlist,
|
||
|
|
*,
|
||
|
|
strict: bool,
|
||
|
|
source: str,
|
||
|
|
log: Callable[[str], None],
|
||
|
|
) -> None:
|
||
|
|
"""Keep only what the allowlist permits.
|
||
|
|
|
||
|
|
Disallowed elements are unwrapped: their children move up, so no text
|
||
|
|
is lost. Comments and processing instructions are dropped. Every
|
||
|
|
removal is reported through log; with strict set, any violation raises
|
||
|
|
instead of changing the page.
|
||
|
|
"""
|
||
|
|
violations: list[str] = []
|
||
|
|
for element in list(tree.getroot().iter()):
|
||
|
|
if not isinstance(element.tag, str): # comment or processing instruction
|
||
|
|
if not strict:
|
||
|
|
_remove(element)
|
||
|
|
continue
|
||
|
|
name = element_name(element.tag)
|
||
|
|
if name not in allowlist:
|
||
|
|
violations.append(f"{source}: <{name}> is not in the Mews subset")
|
||
|
|
if not strict:
|
||
|
|
_unwrap(element)
|
||
|
|
continue
|
||
|
|
allowed = allowlist[name]
|
||
|
|
for key in list(element.attrib):
|
||
|
|
if attribute_name(key) not in allowed:
|
||
|
|
violations.append(
|
||
|
|
f"{source}: {name} {attribute_name(key)} is not in the Mews subset"
|
||
|
|
)
|
||
|
|
if not strict:
|
||
|
|
del element.attrib[key]
|
||
|
|
if not violations:
|
||
|
|
return
|
||
|
|
for message in violations:
|
||
|
|
log(message)
|
||
|
|
if strict:
|
||
|
|
message = "content outside the Mews subset; refusing to write it:\n "
|
||
|
|
raise MewsSubsetError(message + "\n ".join(violations))
|
||
|
|
|
||
|
|
|
||
|
|
def _unwrap(element: Element) -> None:
|
||
|
|
"""Replace an element with its children, keeping the text flow."""
|
||
|
|
parent = element.getparent()
|
||
|
|
index = parent.index(element)
|
||
|
|
children = list(element)
|
||
|
|
if children:
|
||
|
|
children[0].text = (element.text or "") + (children[0].text or "")
|
||
|
|
children[-1].tail = (children[-1].tail or "") + (element.tail or "")
|
||
|
|
for child in children:
|
||
|
|
parent.insert(index, child)
|
||
|
|
index += 1
|
||
|
|
parent.remove(element)
|
||
|
|
return
|
||
|
|
previous = element.getprevious()
|
||
|
|
text = (element.text or "") + (element.tail or "")
|
||
|
|
if previous is not None:
|
||
|
|
previous.tail = (previous.tail or "") + text
|
||
|
|
else:
|
||
|
|
parent.text = (parent.text or "") + text
|
||
|
|
parent.remove(element)
|
||
|
|
|
||
|
|
|
||
|
|
def _remove(element: Element) -> None:
|
||
|
|
"""Drop a comment or processing instruction, keeping the text flow."""
|
||
|
|
parent = element.getparent()
|
||
|
|
previous = element.getprevious()
|
||
|
|
if previous is not None:
|
||
|
|
previous.tail = (previous.tail or "") + (element.tail or "")
|
||
|
|
else:
|
||
|
|
parent.text = (parent.text or "") + (element.tail or "")
|
||
|
|
parent.remove(element)
|