| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225 |
- from __future__ import annotations
- import codecs
- from typing import Any
- from .core import IDNAError, _max_domain_length, _unicode_dots_re, alabel, decode, encode, ulabel
- class Codec(codecs.Codec):
- """Stateless IDNA 2008 codec.
- Implements the :class:`codecs.Codec` protocol so that the whole-domain
- encoder (:func:`idna.encode`) and decoder (:func:`idna.decode`) are
- accessible through the standard codec machinery as ``"idna2008"``.
- Only the ``"strict"`` error handler is supported; any other handler
- raises :exc:`~idna.IDNAError`.
- """
- def encode(self, data: str, errors: str = "strict") -> tuple[bytes, int]: # ty: ignore[invalid-method-override]
- if errors != "strict":
- raise IDNAError(f'Unsupported error handling "{errors}"', code="unsupported_errors")
- if not data:
- return b"", 0
- return encode(data), len(data)
- def decode(self, data: bytes, errors: str = "strict") -> tuple[str, int]: # ty: ignore[invalid-method-override]
- if errors != "strict":
- raise IDNAError(f'Unsupported error handling "{errors}"', code="unsupported_errors")
- if not data:
- return "", 0
- return decode(data), len(data)
- class IncrementalEncoder(codecs.BufferedIncrementalEncoder):
- """Incremental IDNA 2008 encoder.
- Buffers a partial trailing label across calls until either the next
- label separator is seen or ``final=True``, so that streamed input is
- encoded one whole label at a time. Any of the four Unicode label
- separators (``U+002E``, ``U+3002``, ``U+FF0E``, ``U+FF61``) ends a
- label; the result always uses ``U+002E`` as the separator.
- The 253-octet domain length limit (254 with a trailing dot) that
- :func:`idna.encode` enforces is applied to the accumulated output, so
- that streaming a name and encoding it in one shot either both succeed
- with the same result or both raise :exc:`~idna.IDNAError`.
- Only the ``"strict"`` error handler is supported.
- """
- def __init__(self, errors: str = "strict") -> None:
- super().__init__(errors)
- self._emitted = 0 # octets returned so far
- self._trailing_dot = False # whether the output so far ends with "."
- def reset(self) -> None:
- super().reset()
- self._emitted = 0
- self._trailing_dot = False
- def getstate(self) -> Any:
- if not self.buffer and not self._emitted:
- return 0
- return (self.buffer, self._emitted, self._trailing_dot)
- def setstate(self, state: Any) -> None:
- if state:
- self.buffer, self._emitted, self._trailing_dot = state
- else:
- self.reset()
- def _buffer_encode(self, data: str, errors: str, final: bool) -> tuple[bytes, int]: # ty: ignore[invalid-method-override]
- if errors != "strict":
- raise IDNAError(f'Unsupported error handling "{errors}"', code="unsupported_errors")
- result_bytes = b""
- size = 0
- if data:
- labels = _unicode_dots_re.split(data)
- trailing_dot = b""
- if labels:
- if not labels[-1]:
- trailing_dot = b"."
- del labels[-1]
- elif not final:
- # Keep potentially unfinished label until the next call
- del labels[-1]
- if labels:
- trailing_dot = b"."
- result = []
- for label in labels:
- result.append(alabel(label))
- if size:
- size += 1
- size += len(label)
- result_bytes = b".".join(result) + trailing_dot
- size += len(trailing_dot)
- # Mirror encode(): the whole name may not exceed 253 octets, or 254
- # when it ends with a dot. Until the input is final a trailing dot
- # may still arrive, so only the 254-octet ceiling applies before then.
- self._emitted += len(result_bytes)
- if result_bytes:
- self._trailing_dot = result_bytes.endswith(b".")
- may_end_with_dot = self._trailing_dot or not final
- if self._emitted > _max_domain_length + may_end_with_dot:
- raise IDNAError("Domain too long", code="domain_too_long")
- return result_bytes, size
- class IncrementalDecoder(codecs.BufferedIncrementalDecoder):
- """Incremental IDNA 2008 decoder.
- Buffers a partial trailing label across calls until either the next
- label separator is seen or ``final=True``, so that streamed input is
- decoded one whole label at a time.
- The 254-octet input length limit that :func:`idna.decode` enforces is
- applied to the accumulated input, so that streaming a name and decoding
- it in one shot either both succeed with the same result or both raise
- :exc:`~idna.IDNAError`.
- Only the ``"strict"`` error handler is supported.
- """
- def __init__(self, errors: str = "strict") -> None:
- super().__init__(errors)
- self._consumed = 0 # input octets consumed so far
- def reset(self) -> None:
- super().reset()
- self._consumed = 0
- def getstate(self) -> tuple[bytes, int]:
- return (self.buffer, self._consumed)
- def setstate(self, state: tuple[bytes, int]) -> None:
- self.buffer, self._consumed = state
- def _buffer_decode(self, data: Any, errors: str, final: bool) -> tuple[str, int]: # ty: ignore[invalid-method-override]
- if errors != "strict":
- raise IDNAError(f'Unsupported error handling "{errors}"', code="unsupported_errors")
- if not data:
- return ("", 0)
- if not isinstance(data, str):
- try:
- data = str(data, "ascii")
- except UnicodeDecodeError as err:
- raise IDNAError("Invalid ASCII in A-label", code="invalid_ascii") from err
- # Mirror decode(), which rejects input longer than 254 characters
- # (253 plus a possible trailing dot) before looking at any label.
- # ``data`` is the unconsumed buffer plus the new input, so this is
- # the total seen so far.
- if self._consumed + len(data) > _max_domain_length + 1:
- raise IDNAError("Domain too long", code="domain_too_long")
- labels = _unicode_dots_re.split(data)
- trailing_dot = ""
- if labels:
- if not labels[-1]:
- trailing_dot = "."
- del labels[-1]
- elif not final:
- # Keep potentially unfinished label until the next call
- del labels[-1]
- if labels:
- trailing_dot = "."
- result = []
- size = 0
- for label in labels:
- result.append(ulabel(label))
- if size:
- size += 1
- size += len(label)
- result_str = ".".join(result) + trailing_dot
- size += len(trailing_dot)
- self._consumed += size
- return (result_str, size)
- class StreamWriter(Codec, codecs.StreamWriter):
- pass
- class StreamReader(Codec, codecs.StreamReader):
- pass
- def search_function(name: str) -> codecs.CodecInfo | None:
- """Codec search function registered with :mod:`codecs`.
- Returns a :class:`codecs.CodecInfo` for the ``"idna2008"`` codec name
- so that ``str.encode("idna2008")`` and ``bytes.decode("idna2008")``
- invoke the IDNA 2008 codec defined in this module.
- :param name: The codec name being looked up.
- :returns: A :class:`codecs.CodecInfo` instance if ``name`` is
- ``"idna2008"``, otherwise ``None``.
- """
- if name != "idna2008":
- return None
- return codecs.CodecInfo(
- name=name,
- encode=Codec().encode,
- decode=Codec().decode, # type: ignore
- incrementalencoder=IncrementalEncoder,
- incrementaldecoder=IncrementalDecoder,
- streamwriter=StreamWriter,
- streamreader=StreamReader,
- )
- codecs.register(search_function)
|