_winconsole.py 8.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297
  1. # This module is based on the excellent work by Adam Bartoš who
  2. # provided a lot of what went into the implementation here in
  3. # the discussion to issue1602 in the Python bug tracker.
  4. #
  5. # There are some general differences in regards to how this works
  6. # compared to the original patches as we do not need to patch
  7. # the entire interpreter but just work in our little world of
  8. # echo and prompt.
  9. from __future__ import annotations
  10. import collections.abc as cabc
  11. import io
  12. import sys
  13. import time
  14. import typing as t
  15. from ctypes import Array
  16. from ctypes import byref
  17. from ctypes import c_char
  18. from ctypes import c_char_p
  19. from ctypes import c_int
  20. from ctypes import c_ssize_t
  21. from ctypes import c_ulong
  22. from ctypes import c_void_p
  23. from ctypes import POINTER
  24. from ctypes import py_object
  25. from ctypes import Structure
  26. from ctypes.wintypes import DWORD
  27. from ctypes.wintypes import HANDLE
  28. from ctypes.wintypes import LPCWSTR
  29. from ctypes.wintypes import LPWSTR
  30. from gettext import gettext as _
  31. from ._compat import _NonClosingTextIOWrapper
  32. assert sys.platform == "win32"
  33. import msvcrt # noqa: E402
  34. from ctypes import windll # noqa: E402
  35. from ctypes import WINFUNCTYPE # noqa: E402
  36. c_ssize_p = POINTER(c_ssize_t)
  37. kernel32 = windll.kernel32
  38. GetStdHandle = kernel32.GetStdHandle
  39. ReadConsoleW = kernel32.ReadConsoleW
  40. WriteConsoleW = kernel32.WriteConsoleW
  41. GetConsoleMode = kernel32.GetConsoleMode
  42. GetLastError = kernel32.GetLastError
  43. GetCommandLineW = WINFUNCTYPE(LPWSTR)(("GetCommandLineW", windll.kernel32))
  44. CommandLineToArgvW = WINFUNCTYPE(POINTER(LPWSTR), LPCWSTR, POINTER(c_int))(
  45. ("CommandLineToArgvW", windll.shell32)
  46. )
  47. LocalFree = WINFUNCTYPE(c_void_p, c_void_p)(("LocalFree", windll.kernel32))
  48. STDIN_HANDLE = GetStdHandle(-10)
  49. STDOUT_HANDLE = GetStdHandle(-11)
  50. STDERR_HANDLE = GetStdHandle(-12)
  51. PyBUF_SIMPLE = 0
  52. PyBUF_WRITABLE = 1
  53. ERROR_SUCCESS = 0
  54. ERROR_NOT_ENOUGH_MEMORY = 8
  55. ERROR_OPERATION_ABORTED = 995
  56. STDIN_FILENO = 0
  57. STDOUT_FILENO = 1
  58. STDERR_FILENO = 2
  59. EOF = b"\x1a"
  60. MAX_BYTES_WRITTEN = 32767
  61. if t.TYPE_CHECKING:
  62. try:
  63. # Using `typing_extensions.Buffer` instead of `collections.abc`
  64. # on Windows for some reason does not have `Sized` implemented.
  65. from collections.abc import Buffer # type: ignore
  66. except ImportError:
  67. from typing_extensions import Buffer
  68. try:
  69. from ctypes import pythonapi
  70. except ImportError:
  71. # On PyPy we cannot get buffers so our ability to operate here is
  72. # severely limited.
  73. get_buffer = None
  74. else:
  75. class Py_buffer(Structure):
  76. _fields_ = [ # noqa: RUF012
  77. ("buf", c_void_p),
  78. ("obj", py_object),
  79. ("len", c_ssize_t),
  80. ("itemsize", c_ssize_t),
  81. ("readonly", c_int),
  82. ("ndim", c_int),
  83. ("format", c_char_p),
  84. ("shape", c_ssize_p),
  85. ("strides", c_ssize_p),
  86. ("suboffsets", c_ssize_p),
  87. ("internal", c_void_p),
  88. ]
  89. PyObject_GetBuffer = pythonapi.PyObject_GetBuffer
  90. PyBuffer_Release = pythonapi.PyBuffer_Release
  91. def get_buffer(obj: Buffer, writable: bool = False) -> Array[c_char]:
  92. buf = Py_buffer()
  93. flags: int = PyBUF_WRITABLE if writable else PyBUF_SIMPLE
  94. PyObject_GetBuffer(py_object(obj), byref(buf), flags)
  95. try:
  96. buffer_type = c_char * buf.len
  97. out: Array[c_char] = buffer_type.from_address(buf.buf)
  98. return out
  99. finally:
  100. PyBuffer_Release(byref(buf))
  101. class _WindowsConsoleRawIOBase(io.RawIOBase):
  102. def __init__(self, handle: int | None) -> None:
  103. self.handle = handle
  104. def isatty(self) -> t.Literal[True]:
  105. super().isatty()
  106. return True
  107. class _WindowsConsoleReader(_WindowsConsoleRawIOBase):
  108. def readable(self) -> t.Literal[True]:
  109. return True
  110. def readinto(self, b: Buffer) -> int:
  111. bytes_to_be_read = len(b)
  112. if not bytes_to_be_read:
  113. return 0
  114. elif bytes_to_be_read % 2:
  115. raise ValueError(
  116. "cannot read odd number of bytes from UTF-16-LE encoded console"
  117. )
  118. buffer = get_buffer(b, writable=True)
  119. code_units_to_be_read = bytes_to_be_read // 2
  120. code_units_read = c_ulong()
  121. rv = ReadConsoleW(
  122. HANDLE(self.handle),
  123. buffer,
  124. code_units_to_be_read,
  125. byref(code_units_read),
  126. None,
  127. )
  128. if GetLastError() == ERROR_OPERATION_ABORTED:
  129. # wait for KeyboardInterrupt
  130. time.sleep(0.1)
  131. if not rv:
  132. raise OSError(_("Windows error: {error}").format(error=GetLastError()))
  133. if buffer[0] == EOF:
  134. return 0
  135. return 2 * code_units_read.value
  136. class _WindowsConsoleWriter(_WindowsConsoleRawIOBase):
  137. def writable(self) -> t.Literal[True]:
  138. return True
  139. @staticmethod
  140. def _get_error_message(errno: int) -> str:
  141. if errno == ERROR_SUCCESS:
  142. return "ERROR_SUCCESS"
  143. elif errno == ERROR_NOT_ENOUGH_MEMORY:
  144. return "ERROR_NOT_ENOUGH_MEMORY"
  145. return _("Windows error: {error}").format(error=errno)
  146. def write(self, b: Buffer) -> int:
  147. bytes_to_be_written = len(b)
  148. buf = get_buffer(b)
  149. code_units_to_be_written = min(bytes_to_be_written, MAX_BYTES_WRITTEN) // 2
  150. code_units_written = c_ulong()
  151. WriteConsoleW(
  152. HANDLE(self.handle),
  153. buf,
  154. code_units_to_be_written,
  155. byref(code_units_written),
  156. None,
  157. )
  158. bytes_written = 2 * code_units_written.value
  159. if bytes_written == 0 and bytes_to_be_written > 0:
  160. raise OSError(self._get_error_message(GetLastError()))
  161. return bytes_written
  162. class ConsoleStream:
  163. def __init__(self, text_stream: t.TextIO, byte_stream: t.BinaryIO) -> None:
  164. self._text_stream = text_stream
  165. self.buffer = byte_stream
  166. @property
  167. def name(self) -> str:
  168. return self.buffer.name
  169. def write(self, x: t.AnyStr) -> int:
  170. if isinstance(x, str):
  171. return self._text_stream.write(x)
  172. try:
  173. self.flush()
  174. except Exception:
  175. pass
  176. return self.buffer.write(x)
  177. def writelines(self, lines: cabc.Iterable[t.AnyStr]) -> None:
  178. for line in lines:
  179. self.write(line)
  180. def __getattr__(self, name: str) -> t.Any:
  181. return getattr(self._text_stream, name)
  182. def isatty(self) -> bool:
  183. return self.buffer.isatty()
  184. def __repr__(self) -> str:
  185. return f"<ConsoleStream name={self.name!r} encoding={self.encoding!r}>"
  186. def _get_text_stdin(buffer_stream: t.BinaryIO) -> t.TextIO:
  187. text_stream = _NonClosingTextIOWrapper(
  188. io.BufferedReader(_WindowsConsoleReader(STDIN_HANDLE)),
  189. "utf-16-le",
  190. "strict",
  191. line_buffering=True,
  192. )
  193. return t.cast(t.TextIO, ConsoleStream(text_stream, buffer_stream))
  194. def _get_text_stdout(buffer_stream: t.BinaryIO) -> t.TextIO:
  195. text_stream = _NonClosingTextIOWrapper(
  196. io.BufferedWriter(_WindowsConsoleWriter(STDOUT_HANDLE)),
  197. "utf-16-le",
  198. "strict",
  199. line_buffering=True,
  200. )
  201. return t.cast(t.TextIO, ConsoleStream(text_stream, buffer_stream))
  202. def _get_text_stderr(buffer_stream: t.BinaryIO) -> t.TextIO:
  203. text_stream = _NonClosingTextIOWrapper(
  204. io.BufferedWriter(_WindowsConsoleWriter(STDERR_HANDLE)),
  205. "utf-16-le",
  206. "strict",
  207. line_buffering=True,
  208. )
  209. return t.cast(t.TextIO, ConsoleStream(text_stream, buffer_stream))
  210. _stream_factories: cabc.Mapping[int, t.Callable[[t.BinaryIO], t.TextIO]] = {
  211. 0: _get_text_stdin,
  212. 1: _get_text_stdout,
  213. 2: _get_text_stderr,
  214. }
  215. def _is_console(f: t.TextIO) -> bool:
  216. if not hasattr(f, "fileno"):
  217. return False
  218. try:
  219. fileno = f.fileno()
  220. except (OSError, io.UnsupportedOperation):
  221. return False
  222. handle = msvcrt.get_osfhandle(fileno)
  223. return bool(GetConsoleMode(handle, byref(DWORD())))
  224. def _get_windows_console_stream(
  225. f: t.TextIO, encoding: str | None, errors: str | None
  226. ) -> t.TextIO | None:
  227. if (
  228. get_buffer is None
  229. or encoding not in {"utf-16-le", None}
  230. or errors not in {"strict", None}
  231. or not _is_console(f)
  232. ):
  233. return None
  234. func = _stream_factories.get(f.fileno())
  235. if func is None:
  236. return None
  237. b = getattr(f, "buffer", None)
  238. if b is None:
  239. return None
  240. return func(b)