codec.py 7.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225
  1. from __future__ import annotations
  2. import codecs
  3. from typing import Any
  4. from .core import IDNAError, _max_domain_length, _unicode_dots_re, alabel, decode, encode, ulabel
  5. class Codec(codecs.Codec):
  6. """Stateless IDNA 2008 codec.
  7. Implements the :class:`codecs.Codec` protocol so that the whole-domain
  8. encoder (:func:`idna.encode`) and decoder (:func:`idna.decode`) are
  9. accessible through the standard codec machinery as ``"idna2008"``.
  10. Only the ``"strict"`` error handler is supported; any other handler
  11. raises :exc:`~idna.IDNAError`.
  12. """
  13. def encode(self, data: str, errors: str = "strict") -> tuple[bytes, int]: # ty: ignore[invalid-method-override]
  14. if errors != "strict":
  15. raise IDNAError(f'Unsupported error handling "{errors}"', code="unsupported_errors")
  16. if not data:
  17. return b"", 0
  18. return encode(data), len(data)
  19. def decode(self, data: bytes, errors: str = "strict") -> tuple[str, int]: # ty: ignore[invalid-method-override]
  20. if errors != "strict":
  21. raise IDNAError(f'Unsupported error handling "{errors}"', code="unsupported_errors")
  22. if not data:
  23. return "", 0
  24. return decode(data), len(data)
  25. class IncrementalEncoder(codecs.BufferedIncrementalEncoder):
  26. """Incremental IDNA 2008 encoder.
  27. Buffers a partial trailing label across calls until either the next
  28. label separator is seen or ``final=True``, so that streamed input is
  29. encoded one whole label at a time. Any of the four Unicode label
  30. separators (``U+002E``, ``U+3002``, ``U+FF0E``, ``U+FF61``) ends a
  31. label; the result always uses ``U+002E`` as the separator.
  32. The 253-octet domain length limit (254 with a trailing dot) that
  33. :func:`idna.encode` enforces is applied to the accumulated output, so
  34. that streaming a name and encoding it in one shot either both succeed
  35. with the same result or both raise :exc:`~idna.IDNAError`.
  36. Only the ``"strict"`` error handler is supported.
  37. """
  38. def __init__(self, errors: str = "strict") -> None:
  39. super().__init__(errors)
  40. self._emitted = 0 # octets returned so far
  41. self._trailing_dot = False # whether the output so far ends with "."
  42. def reset(self) -> None:
  43. super().reset()
  44. self._emitted = 0
  45. self._trailing_dot = False
  46. def getstate(self) -> Any:
  47. if not self.buffer and not self._emitted:
  48. return 0
  49. return (self.buffer, self._emitted, self._trailing_dot)
  50. def setstate(self, state: Any) -> None:
  51. if state:
  52. self.buffer, self._emitted, self._trailing_dot = state
  53. else:
  54. self.reset()
  55. def _buffer_encode(self, data: str, errors: str, final: bool) -> tuple[bytes, int]: # ty: ignore[invalid-method-override]
  56. if errors != "strict":
  57. raise IDNAError(f'Unsupported error handling "{errors}"', code="unsupported_errors")
  58. result_bytes = b""
  59. size = 0
  60. if data:
  61. labels = _unicode_dots_re.split(data)
  62. trailing_dot = b""
  63. if labels:
  64. if not labels[-1]:
  65. trailing_dot = b"."
  66. del labels[-1]
  67. elif not final:
  68. # Keep potentially unfinished label until the next call
  69. del labels[-1]
  70. if labels:
  71. trailing_dot = b"."
  72. result = []
  73. for label in labels:
  74. result.append(alabel(label))
  75. if size:
  76. size += 1
  77. size += len(label)
  78. result_bytes = b".".join(result) + trailing_dot
  79. size += len(trailing_dot)
  80. # Mirror encode(): the whole name may not exceed 253 octets, or 254
  81. # when it ends with a dot. Until the input is final a trailing dot
  82. # may still arrive, so only the 254-octet ceiling applies before then.
  83. self._emitted += len(result_bytes)
  84. if result_bytes:
  85. self._trailing_dot = result_bytes.endswith(b".")
  86. may_end_with_dot = self._trailing_dot or not final
  87. if self._emitted > _max_domain_length + may_end_with_dot:
  88. raise IDNAError("Domain too long", code="domain_too_long")
  89. return result_bytes, size
  90. class IncrementalDecoder(codecs.BufferedIncrementalDecoder):
  91. """Incremental IDNA 2008 decoder.
  92. Buffers a partial trailing label across calls until either the next
  93. label separator is seen or ``final=True``, so that streamed input is
  94. decoded one whole label at a time.
  95. The 254-octet input length limit that :func:`idna.decode` enforces is
  96. applied to the accumulated input, so that streaming a name and decoding
  97. it in one shot either both succeed with the same result or both raise
  98. :exc:`~idna.IDNAError`.
  99. Only the ``"strict"`` error handler is supported.
  100. """
  101. def __init__(self, errors: str = "strict") -> None:
  102. super().__init__(errors)
  103. self._consumed = 0 # input octets consumed so far
  104. def reset(self) -> None:
  105. super().reset()
  106. self._consumed = 0
  107. def getstate(self) -> tuple[bytes, int]:
  108. return (self.buffer, self._consumed)
  109. def setstate(self, state: tuple[bytes, int]) -> None:
  110. self.buffer, self._consumed = state
  111. def _buffer_decode(self, data: Any, errors: str, final: bool) -> tuple[str, int]: # ty: ignore[invalid-method-override]
  112. if errors != "strict":
  113. raise IDNAError(f'Unsupported error handling "{errors}"', code="unsupported_errors")
  114. if not data:
  115. return ("", 0)
  116. if not isinstance(data, str):
  117. try:
  118. data = str(data, "ascii")
  119. except UnicodeDecodeError as err:
  120. raise IDNAError("Invalid ASCII in A-label", code="invalid_ascii") from err
  121. # Mirror decode(), which rejects input longer than 254 characters
  122. # (253 plus a possible trailing dot) before looking at any label.
  123. # ``data`` is the unconsumed buffer plus the new input, so this is
  124. # the total seen so far.
  125. if self._consumed + len(data) > _max_domain_length + 1:
  126. raise IDNAError("Domain too long", code="domain_too_long")
  127. labels = _unicode_dots_re.split(data)
  128. trailing_dot = ""
  129. if labels:
  130. if not labels[-1]:
  131. trailing_dot = "."
  132. del labels[-1]
  133. elif not final:
  134. # Keep potentially unfinished label until the next call
  135. del labels[-1]
  136. if labels:
  137. trailing_dot = "."
  138. result = []
  139. size = 0
  140. for label in labels:
  141. result.append(ulabel(label))
  142. if size:
  143. size += 1
  144. size += len(label)
  145. result_str = ".".join(result) + trailing_dot
  146. size += len(trailing_dot)
  147. self._consumed += size
  148. return (result_str, size)
  149. class StreamWriter(Codec, codecs.StreamWriter):
  150. pass
  151. class StreamReader(Codec, codecs.StreamReader):
  152. pass
  153. def search_function(name: str) -> codecs.CodecInfo | None:
  154. """Codec search function registered with :mod:`codecs`.
  155. Returns a :class:`codecs.CodecInfo` for the ``"idna2008"`` codec name
  156. so that ``str.encode("idna2008")`` and ``bytes.decode("idna2008")``
  157. invoke the IDNA 2008 codec defined in this module.
  158. :param name: The codec name being looked up.
  159. :returns: A :class:`codecs.CodecInfo` instance if ``name`` is
  160. ``"idna2008"``, otherwise ``None``.
  161. """
  162. if name != "idna2008":
  163. return None
  164. return codecs.CodecInfo(
  165. name=name,
  166. encode=Codec().encode,
  167. decode=Codec().decode, # type: ignore
  168. incrementalencoder=IncrementalEncoder,
  169. incrementaldecoder=IncrementalDecoder,
  170. streamwriter=StreamWriter,
  171. streamreader=StreamReader,
  172. )
  173. codecs.register(search_function)