_compat.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590
  1. from __future__ import annotations
  2. import codecs
  3. import collections.abc as cabc
  4. import io
  5. import os
  6. import re
  7. import sys
  8. import typing as t
  9. from types import TracebackType
  10. from weakref import WeakKeyDictionary
  11. CYGWIN = sys.platform.startswith("cygwin")
  12. WIN = sys.platform.startswith("win")
  13. # One CSI escape sequence per the ECMA-48 grammar: parameter bytes (0x30-0x3F),
  14. # intermediate bytes (0x20-0x2F), then a final byte (0x40-0x7E). Broader than the
  15. # SGR codes Click emits, so foreign sequences (colon-delimited true-color, mouse
  16. # reporting) are stripped too.
  17. _ansi_re = re.compile(r"\033\[[0-?]*[ -/]*[@-~]")
  18. def _make_text_stream(
  19. stream: t.BinaryIO,
  20. encoding: str | None,
  21. errors: str | None,
  22. force_readable: bool = False,
  23. force_writable: bool = False,
  24. ) -> t.TextIO:
  25. if encoding is None:
  26. encoding = get_best_encoding(stream)
  27. if errors is None:
  28. errors = "replace"
  29. return _NonClosingTextIOWrapper(
  30. stream,
  31. encoding,
  32. errors,
  33. line_buffering=True,
  34. force_readable=force_readable,
  35. force_writable=force_writable,
  36. )
  37. def is_ascii_encoding(encoding: str) -> bool:
  38. """Checks if a given encoding is ascii."""
  39. try:
  40. return codecs.lookup(encoding).name == "ascii"
  41. except LookupError:
  42. return False
  43. def get_best_encoding(stream: t.IO[t.Any]) -> str:
  44. """Returns the default stream encoding if not found."""
  45. rv = getattr(stream, "encoding", None) or sys.getdefaultencoding()
  46. if is_ascii_encoding(rv):
  47. return "utf-8"
  48. return rv
  49. class _NonClosingTextIOWrapper(io.TextIOWrapper):
  50. def __init__(
  51. self,
  52. stream: t.BinaryIO,
  53. encoding: str | None,
  54. errors: str | None,
  55. force_readable: bool = False,
  56. force_writable: bool = False,
  57. **extra: t.Any,
  58. ) -> None:
  59. self._stream = stream = t.cast(
  60. t.BinaryIO, _FixupStream(stream, force_readable, force_writable)
  61. )
  62. super().__init__(stream, encoding, errors, **extra)
  63. def __del__(self) -> None:
  64. try:
  65. self.detach()
  66. except Exception:
  67. pass
  68. def isatty(self) -> bool:
  69. # https://bitbucket.org/pypy/pypy/issue/1803
  70. return self._stream.isatty()
  71. class _FixupStream:
  72. """The new io interface needs more from streams than streams
  73. traditionally implement. As such, this fix-up code is necessary in
  74. some circumstances.
  75. The forcing of readable and writable flags are there because some tools
  76. put badly patched objects on sys (one such offender are certain version
  77. of jupyter notebook).
  78. """
  79. def __init__(
  80. self,
  81. stream: t.BinaryIO,
  82. force_readable: bool = False,
  83. force_writable: bool = False,
  84. ):
  85. self._stream = stream
  86. self._force_readable = force_readable
  87. self._force_writable = force_writable
  88. def __getattr__(self, name: str) -> t.Any:
  89. return getattr(self._stream, name)
  90. def read1(self, size: int) -> bytes:
  91. f = getattr(self._stream, "read1", None)
  92. if f is not None:
  93. return t.cast(bytes, f(size))
  94. return self._stream.read(size)
  95. def readable(self) -> bool:
  96. if self._force_readable:
  97. return True
  98. x = getattr(self._stream, "readable", None)
  99. if x is not None:
  100. return t.cast(bool, x())
  101. try:
  102. self._stream.read(0)
  103. except Exception:
  104. return False
  105. return True
  106. def writable(self) -> bool:
  107. if self._force_writable:
  108. return True
  109. x = getattr(self._stream, "writable", None)
  110. if x is not None:
  111. return t.cast(bool, x())
  112. try:
  113. self._stream.write(b"")
  114. except Exception:
  115. try:
  116. self._stream.write(b"")
  117. except Exception:
  118. return False
  119. return True
  120. def seekable(self) -> bool:
  121. x = getattr(self._stream, "seekable", None)
  122. if x is not None:
  123. return t.cast(bool, x())
  124. try:
  125. self._stream.seek(self._stream.tell())
  126. except Exception:
  127. return False
  128. return True
  129. def _is_binary_reader(stream: t.IO[t.Any], default: bool = False) -> bool:
  130. try:
  131. return isinstance(stream.read(0), bytes)
  132. except Exception:
  133. return default
  134. # This happens in some cases where the stream was already
  135. # closed. In this case, we assume the default.
  136. def _is_binary_writer(stream: t.IO[t.Any], default: bool = False) -> bool:
  137. try:
  138. stream.write(b"")
  139. except Exception:
  140. try:
  141. stream.write("")
  142. return False
  143. except Exception:
  144. pass
  145. return default
  146. return True
  147. def _find_binary_reader(stream: t.IO[t.Any]) -> t.BinaryIO | None:
  148. # We need to figure out if the given stream is already binary.
  149. # This can happen because the official docs recommend detaching
  150. # the streams to get binary streams. Some code might do this, so
  151. # we need to deal with this case explicitly.
  152. if _is_binary_reader(stream, False):
  153. return t.cast(t.BinaryIO, stream)
  154. buf = getattr(stream, "buffer", None)
  155. # Same situation here; this time we assume that the buffer is
  156. # actually binary in case it's closed.
  157. if buf is not None and _is_binary_reader(buf, True):
  158. return t.cast(t.BinaryIO, buf)
  159. return None
  160. def _find_binary_writer(stream: t.IO[t.Any]) -> t.BinaryIO | None:
  161. # We need to figure out if the given stream is already binary.
  162. # This can happen because the official docs recommend detaching
  163. # the streams to get binary streams. Some code might do this, so
  164. # we need to deal with this case explicitly.
  165. if _is_binary_writer(stream, False):
  166. return t.cast(t.BinaryIO, stream)
  167. buf = getattr(stream, "buffer", None)
  168. # Same situation here; this time we assume that the buffer is
  169. # actually binary in case it's closed.
  170. if buf is not None and _is_binary_writer(buf, True):
  171. return t.cast(t.BinaryIO, buf)
  172. return None
  173. def _stream_is_misconfigured(stream: t.TextIO) -> bool:
  174. """A stream is misconfigured if its encoding is ASCII."""
  175. # If the stream does not have an encoding set, we assume it's set
  176. # to ASCII. This appears to happen in certain unittest
  177. # environments. It's not quite clear what the correct behavior is
  178. # but this at least will force Click to recover somehow.
  179. return is_ascii_encoding(getattr(stream, "encoding", None) or "ascii")
  180. def _is_compat_stream_attr(stream: t.TextIO, attr: str, value: str | None) -> bool:
  181. """A stream attribute is compatible if it is equal to the
  182. desired value or the desired value is unset and the attribute
  183. has a value.
  184. """
  185. stream_value = getattr(stream, attr, None)
  186. return stream_value == value or (value is None and stream_value is not None)
  187. def _is_compatible_text_stream(
  188. stream: t.TextIO, encoding: str | None, errors: str | None
  189. ) -> bool:
  190. """Check if a stream's encoding and errors attributes are
  191. compatible with the desired values.
  192. """
  193. return _is_compat_stream_attr(
  194. stream, "encoding", encoding
  195. ) and _is_compat_stream_attr(stream, "errors", errors)
  196. def _force_correct_text_stream(
  197. text_stream: t.IO[t.Any],
  198. encoding: str | None,
  199. errors: str | None,
  200. is_binary: t.Callable[[t.IO[t.Any], bool], bool],
  201. find_binary: t.Callable[[t.IO[t.Any]], t.BinaryIO | None],
  202. force_readable: bool = False,
  203. force_writable: bool = False,
  204. ) -> t.TextIO:
  205. if is_binary(text_stream, False):
  206. binary_reader = t.cast(t.BinaryIO, text_stream)
  207. else:
  208. text_stream = t.cast(t.TextIO, text_stream)
  209. # If the stream looks compatible, and won't default to a
  210. # misconfigured ascii encoding, return it as-is.
  211. if _is_compatible_text_stream(text_stream, encoding, errors) and not (
  212. encoding is None and _stream_is_misconfigured(text_stream)
  213. ):
  214. return text_stream
  215. # Otherwise, get the underlying binary reader.
  216. possible_binary_reader = find_binary(text_stream)
  217. # If that's not possible, silently use the original reader
  218. # and get mojibake instead of exceptions.
  219. if possible_binary_reader is None:
  220. return text_stream
  221. binary_reader = possible_binary_reader
  222. # Default errors to replace instead of strict in order to get
  223. # something that works.
  224. if errors is None:
  225. errors = "replace"
  226. # Wrap the binary stream in a text stream with the correct
  227. # encoding parameters.
  228. return _make_text_stream(
  229. binary_reader,
  230. encoding,
  231. errors,
  232. force_readable=force_readable,
  233. force_writable=force_writable,
  234. )
  235. def _force_correct_text_reader(
  236. text_reader: t.IO[t.Any],
  237. encoding: str | None,
  238. errors: str | None,
  239. force_readable: bool = False,
  240. ) -> t.TextIO:
  241. return _force_correct_text_stream(
  242. text_reader,
  243. encoding,
  244. errors,
  245. _is_binary_reader,
  246. _find_binary_reader,
  247. force_readable=force_readable,
  248. )
  249. def _force_correct_text_writer(
  250. text_writer: t.IO[t.Any],
  251. encoding: str | None,
  252. errors: str | None,
  253. force_writable: bool = False,
  254. ) -> t.TextIO:
  255. return _force_correct_text_stream(
  256. text_writer,
  257. encoding,
  258. errors,
  259. _is_binary_writer,
  260. _find_binary_writer,
  261. force_writable=force_writable,
  262. )
  263. def get_binary_stdin() -> t.BinaryIO:
  264. reader = _find_binary_reader(sys.stdin)
  265. if reader is None:
  266. raise RuntimeError("Was not able to determine binary stream for sys.stdin.")
  267. return reader
  268. def get_binary_stdout() -> t.BinaryIO:
  269. writer = _find_binary_writer(sys.stdout)
  270. if writer is None:
  271. raise RuntimeError("Was not able to determine binary stream for sys.stdout.")
  272. return writer
  273. def get_binary_stderr() -> t.BinaryIO:
  274. writer = _find_binary_writer(sys.stderr)
  275. if writer is None:
  276. raise RuntimeError("Was not able to determine binary stream for sys.stderr.")
  277. return writer
  278. def get_text_stdin(encoding: str | None = None, errors: str | None = None) -> t.TextIO:
  279. rv = _get_windows_console_stream(sys.stdin, encoding, errors)
  280. if rv is not None:
  281. return rv
  282. return _force_correct_text_reader(sys.stdin, encoding, errors, force_readable=True)
  283. def get_text_stdout(encoding: str | None = None, errors: str | None = None) -> t.TextIO:
  284. rv = _get_windows_console_stream(sys.stdout, encoding, errors)
  285. if rv is not None:
  286. return rv
  287. return _force_correct_text_writer(sys.stdout, encoding, errors, force_writable=True)
  288. def get_text_stderr(encoding: str | None = None, errors: str | None = None) -> t.TextIO:
  289. rv = _get_windows_console_stream(sys.stderr, encoding, errors)
  290. if rv is not None:
  291. return rv
  292. return _force_correct_text_writer(sys.stderr, encoding, errors, force_writable=True)
  293. def _wrap_io_open(
  294. file: str | os.PathLike[str] | int,
  295. mode: str,
  296. encoding: str | None,
  297. errors: str | None,
  298. ) -> t.IO[t.Any]:
  299. """Handles not passing ``encoding`` and ``errors`` in binary mode."""
  300. if "b" in mode:
  301. return open(file, mode)
  302. return open(file, mode, encoding=encoding, errors=errors)
  303. def open_stream(
  304. filename: str | os.PathLike[str],
  305. mode: str = "r",
  306. encoding: str | None = None,
  307. errors: str | None = "strict",
  308. atomic: bool = False,
  309. ) -> tuple[t.IO[t.Any], bool]:
  310. binary = "b" in mode
  311. filename = os.fspath(filename)
  312. # Standard streams first. These are simple because they ignore the
  313. # atomic flag. Use fsdecode to handle Path("-").
  314. if os.fsdecode(filename) == "-":
  315. if any(m in mode for m in ["w", "a", "x"]):
  316. if binary:
  317. return get_binary_stdout(), False
  318. return get_text_stdout(encoding=encoding, errors=errors), False
  319. if binary:
  320. return get_binary_stdin(), False
  321. return get_text_stdin(encoding=encoding, errors=errors), False
  322. # Non-atomic writes directly go out through the regular open functions.
  323. if not atomic:
  324. return _wrap_io_open(filename, mode, encoding, errors), True
  325. # Some usability stuff for atomic writes
  326. if "a" in mode:
  327. raise ValueError(
  328. "Appending to an existing file is not supported, because that"
  329. " would involve an expensive `copy`-operation to a temporary"
  330. " file. Open the file in normal `w`-mode and copy explicitly"
  331. " if that's what you're after."
  332. )
  333. if "x" in mode:
  334. raise ValueError("Use the `overwrite`-parameter instead.")
  335. if "w" not in mode:
  336. raise ValueError("Atomic writes only make sense with `w`-mode.")
  337. # Atomic writes are more complicated. They work by opening a file
  338. # as a proxy in the same folder and then using the fdopen
  339. # functionality to wrap it in a Python file. Then we wrap it in an
  340. # atomic file that moves the file over on close.
  341. import errno
  342. import random
  343. try:
  344. perm: int | None = os.stat(filename).st_mode
  345. except OSError:
  346. perm = None
  347. flags = os.O_RDWR | os.O_CREAT | os.O_EXCL
  348. if binary:
  349. flags |= getattr(os, "O_BINARY", 0)
  350. while True:
  351. tmp_filename = os.path.join(
  352. os.path.dirname(filename),
  353. f".__atomic-write{random.randrange(1 << 32):08x}",
  354. )
  355. try:
  356. fd = os.open(tmp_filename, flags, 0o666 if perm is None else perm)
  357. break
  358. except OSError as e:
  359. if e.errno == errno.EEXIST or (
  360. os.name == "nt"
  361. and e.errno == errno.EACCES
  362. and os.path.isdir(e.filename)
  363. and os.access(e.filename, os.W_OK)
  364. ):
  365. continue
  366. raise
  367. if perm is not None:
  368. os.chmod(tmp_filename, perm) # in case perm includes bits in umask
  369. f = _wrap_io_open(fd, mode, encoding, errors)
  370. af = _AtomicFile(f, tmp_filename, os.path.realpath(filename))
  371. return t.cast(t.IO[t.Any], af), True
  372. class _AtomicFile:
  373. def __init__(self, f: t.IO[t.Any], tmp_filename: str, real_filename: str) -> None:
  374. self._f = f
  375. self._tmp_filename = tmp_filename
  376. self._real_filename = real_filename
  377. self.closed = False
  378. @property
  379. def name(self) -> str:
  380. return self._real_filename
  381. def close(self, delete: bool = False) -> None:
  382. if self.closed:
  383. return
  384. self._f.close()
  385. os.replace(self._tmp_filename, self._real_filename)
  386. self.closed = True
  387. def __getattr__(self, name: str) -> t.Any:
  388. return getattr(self._f, name)
  389. def __enter__(self) -> _AtomicFile:
  390. return self
  391. def __exit__(
  392. self,
  393. exc_type: type[BaseException] | None,
  394. exc_value: BaseException | None,
  395. tb: TracebackType | None,
  396. ) -> None:
  397. self.close(delete=exc_type is not None)
  398. def __repr__(self) -> str:
  399. return repr(self._f)
  400. def strip_ansi(value: str) -> str:
  401. return _ansi_re.sub("", value)
  402. def _is_jupyter_kernel_output(stream: t.IO[t.Any]) -> bool:
  403. while isinstance(stream, (_FixupStream, _NonClosingTextIOWrapper)):
  404. stream = stream._stream
  405. return stream.__class__.__module__.startswith("ipykernel.")
  406. def should_strip_ansi(
  407. stream: t.IO[t.Any] | None = None, color: bool | None = None
  408. ) -> bool:
  409. if color is None:
  410. if stream is None:
  411. stream = sys.stdin
  412. elif hasattr(stream, "color"):
  413. # ._termui_impl._PagerWriter handles stripping ansi itself,
  414. # so we don't need to strip it here
  415. return False
  416. return not isatty(stream) and not _is_jupyter_kernel_output(stream)
  417. return not color
  418. # Double check is needed so mypy does not analyze this on Linux.
  419. if sys.platform.startswith("win") and WIN:
  420. from ._winconsole import _get_windows_console_stream
  421. def _get_argv_encoding() -> str:
  422. import locale
  423. return locale.getpreferredencoding()
  424. else:
  425. def _get_argv_encoding() -> str:
  426. return getattr(sys.stdin, "encoding", None) or sys.getfilesystemencoding()
  427. def _get_windows_console_stream(
  428. f: t.TextIO, encoding: str | None, errors: str | None
  429. ) -> t.TextIO | None:
  430. return None
  431. def term_len(x: str) -> int:
  432. return len(strip_ansi(x))
  433. def isatty(stream: t.IO[t.Any]) -> bool:
  434. try:
  435. return stream.isatty()
  436. except Exception:
  437. return False
  438. def _make_cached_stream_func(
  439. src_func: t.Callable[[], t.TextIO | None],
  440. wrapper_func: t.Callable[[], t.TextIO],
  441. ) -> t.Callable[[], t.TextIO | None]:
  442. cache: cabc.MutableMapping[t.TextIO, t.TextIO] = WeakKeyDictionary()
  443. def func() -> t.TextIO | None:
  444. stream = src_func()
  445. if stream is None:
  446. return None
  447. try:
  448. rv = cache.get(stream)
  449. except Exception:
  450. rv = None
  451. if rv is not None:
  452. return rv
  453. rv = wrapper_func()
  454. try:
  455. cache[stream] = rv
  456. except Exception:
  457. pass
  458. return rv
  459. return func
  460. _default_text_stdin = _make_cached_stream_func(lambda: sys.stdin, get_text_stdin)
  461. _default_text_stdout = _make_cached_stream_func(lambda: sys.stdout, get_text_stdout)
  462. _default_text_stderr = _make_cached_stream_func(lambda: sys.stderr, get_text_stderr)
  463. binary_streams: cabc.Mapping[str, t.Callable[[], t.BinaryIO]] = {
  464. "stdin": get_binary_stdin,
  465. "stdout": get_binary_stdout,
  466. "stderr": get_binary_stderr,
  467. }
  468. text_streams: cabc.Mapping[str, t.Callable[[str | None, str | None], t.TextIO]] = {
  469. "stdin": get_text_stdin,
  470. "stdout": get_text_stdout,
  471. "stderr": get_text_stderr,
  472. }