xmlwriter.py 4.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130
  1. from __future__ import annotations
  2. import codecs
  3. from typing import IO, TYPE_CHECKING, Dict, Iterable, List, Optional, Tuple
  4. from xml.sax.saxutils import escape, quoteattr
  5. from rdflib.term import URIRef
  6. if TYPE_CHECKING:
  7. from rdflib.namespace import Namespace, NamespaceManager
  8. __all__ = ["XMLWriter"]
  9. ESCAPE_ENTITIES = {"\r": "
"}
  10. class XMLWriter:
  11. """A simple XML writer that writes to a stream."""
  12. def __init__(
  13. self,
  14. stream: IO[bytes],
  15. namespace_manager: NamespaceManager,
  16. encoding: Optional[str] = None,
  17. decl: int = 1,
  18. extra_ns: Optional[Dict[str, Namespace]] = None,
  19. ):
  20. encoding = encoding or "utf-8"
  21. encoder, decoder, stream_reader, stream_writer = codecs.lookup(encoding)
  22. # NOTE on type ignores: this is mainly because the variable is being re-used.
  23. # type error: Incompatible types in assignment (expression has type "StreamWriter", variable has type "IO[bytes]")
  24. self.stream = stream = stream_writer(stream) # type: ignore[assignment]
  25. if decl:
  26. # type error: No overload variant of "write" of "IO" matches argument type "str"
  27. stream.write('<?xml version="1.0" encoding="%s"?>' % encoding) # type: ignore[call-overload]
  28. self.element_stack: List[str] = []
  29. self.nm = namespace_manager
  30. self.extra_ns = extra_ns or {}
  31. self.closed = True
  32. def __get_indent(self) -> str:
  33. return " " * len(self.element_stack)
  34. indent = property(__get_indent)
  35. def __close_start_tag(self) -> None:
  36. if not self.closed: # TODO:
  37. self.closed = True
  38. self.stream.write(">")
  39. def push(self, uri: str) -> None:
  40. self.__close_start_tag()
  41. write = self.stream.write
  42. write("\n")
  43. write(self.indent)
  44. write("<%s" % self.qname(uri))
  45. self.element_stack.append(uri)
  46. self.closed = False
  47. self.parent = False
  48. def pop(self, uri: Optional[str] = None) -> None:
  49. top = self.element_stack.pop()
  50. if uri:
  51. assert uri == top
  52. write = self.stream.write
  53. if not self.closed:
  54. self.closed = True
  55. write("/>")
  56. else:
  57. if self.parent:
  58. write("\n")
  59. write(self.indent)
  60. write("</%s>" % self.qname(top))
  61. self.parent = True
  62. def element(
  63. self, uri: str, content: str, attributes: Dict[URIRef, str] = {}
  64. ) -> None:
  65. """Utility method for adding a complete simple element"""
  66. self.push(uri)
  67. for k, v in attributes.items():
  68. self.attribute(k, v)
  69. self.text(content)
  70. self.pop()
  71. def namespaces(self, namespaces: Iterable[Tuple[str, str]] = None) -> None:
  72. if not namespaces:
  73. namespaces = self.nm.namespaces()
  74. write = self.stream.write
  75. write("\n")
  76. for prefix, namespace in namespaces:
  77. if prefix:
  78. write(' xmlns:%s="%s"\n' % (prefix, namespace))
  79. # Allow user-provided namespace bindings to prevail
  80. elif prefix not in self.extra_ns:
  81. write(' xmlns="%s"\n' % namespace)
  82. for prefix, namespace in self.extra_ns.items():
  83. if prefix:
  84. write(' xmlns:%s="%s"\n' % (prefix, namespace))
  85. else:
  86. write(' xmlns="%s"\n' % namespace)
  87. def attribute(self, uri: str, value: str) -> None:
  88. write = self.stream.write
  89. write(" %s=%s" % (self.qname(uri), quoteattr(value)))
  90. def text(self, text: str) -> None:
  91. self.__close_start_tag()
  92. if "<" in text and ">" in text and "]]>" not in text:
  93. self.stream.write("<![CDATA[")
  94. self.stream.write(text)
  95. self.stream.write("]]>")
  96. else:
  97. self.stream.write(escape(text, ESCAPE_ENTITIES))
  98. def qname(self, uri: str) -> str:
  99. """Compute qname for a uri using our extra namespaces,
  100. or the given namespace manager"""
  101. for pre, ns in self.extra_ns.items():
  102. if uri.startswith(ns):
  103. if pre != "":
  104. return ":".join([pre, uri[len(ns) :]])
  105. else:
  106. return uri[len(ns) :]
  107. return self.nm.qname_strict(uri)