"""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 declarations, attributes from 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"", re.DOTALL ) _ATTLIST = re.compile(r"", re.DOTALL) _ELEMENT = re.compile(r" 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 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)