diff --git a/aexpect/client.py b/aexpect/client.py index a369046..e2d7752 100644 --- a/aexpect/client.py +++ b/aexpect/client.py @@ -27,6 +27,7 @@ import subprocess import threading import time +from codecs import getincrementaldecoder from aexpect.exceptions import ( ExpectError, @@ -56,6 +57,9 @@ LOG = logging.getLogger(__name__) +# Buffer size in byte for pipe reads +READ_BUFFER_SIZE = 1024 + def kill_tail_threads(): """ @@ -731,6 +735,8 @@ def _print_line(text): poller = select.poll() poller.register(tail_pipe, select.POLLIN) bfr = "" + decoder_class = getincrementaldecoder(self.encoding) + decoder = decoder_class(errors="ignore") while True: if _THREAD_KILL_REQUESTED.is_set(): try: @@ -745,10 +751,10 @@ def _print_line(text): break if poll_status: # Some data is available; read it - new_data = os.read(tail_pipe, 1024) - if not new_data: + new_bytes = os.read(tail_pipe, READ_BUFFER_SIZE) + if not new_bytes: break - new_data = new_data.decode(self.encoding, "ignore") + new_data = decoder.decode(input=new_bytes) if not new_data: # all chars were ignored, skip round continue bfr += new_data @@ -901,23 +907,23 @@ def _read_nonblocking(self, internal_timeout=None, timeout=None): expect_pipe = self._get_fd("expect") poller = select.poll() poller.register(expect_pipe, select.POLLIN) - data = "" + data = b"" read = 0 while True: try: poll_status = poller.poll(internal_timeout) except select.error: - return read, data + return read, data.decode(self.encoding, "ignore") if poll_status: - raw_data = os.read(expect_pipe, 1024) + raw_data = os.read(expect_pipe, READ_BUFFER_SIZE) if not raw_data: - return read, data + return read, data.decode(self.encoding, "ignore") read += len(raw_data) - data += raw_data.decode(self.encoding, "ignore") + data += raw_data else: - return read, data + return read, data.decode(self.encoding, "ignore") if end_time and time.monotonic() > end_time: - return read, data + return read, data.decode(self.encoding, "ignore") def read_nonblocking(self, internal_timeout=None, timeout=None): """ diff --git a/tests/test_client.py b/tests/test_client.py index 4a7cafb..1a2c33f 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -18,6 +18,7 @@ import random import string import sys +import time import unittest from aexpect import client @@ -194,5 +195,88 @@ def get_proc_fds(): ) +class EncodingTest(unittest.TestCase): + + TEXT = "嗨😀" + MAX_OFFSET = 10 + + def _multibyte_write_cmd(self, offset, count=1): + """Build a Python command that writes multibyte text to stdout.""" + encoded = self.TEXT.encode("utf-8") + reps = 1024 // len(encoded) + 1 + writes = "; ".join(["f.write(t); f.flush()"] * count) + return ( + f"import os,sys; t=b' '*{offset}+{encoded!r}*{reps}+b'\\n'; " + f"f=os.fdopen(sys.stdout.fileno(),'wb',closefd=False); {writes}" + ) + + @unittest.skipUnless(os.name == "posix", "Unix/Linux/macOS only") + def test_shell(self): + """Test multibyte decoding in ShellSession across buffer boundaries.""" + sess = client.ShellSession("/bin/sh") + sess.cmd_output("echo init") + lengths = [] + for offset in range(self.MAX_OFFSET): + cmd = self._multibyte_write_cmd(offset) + result = sess.cmd_output(f'{sys.executable} -c "{cmd}"').lstrip() + self.assertTrue( + result.startswith(self.TEXT), + f"offset {offset}: unexpected start: {result[:20]!r}", + ) + lengths.append(len(result)) + sess.close() + self.assertTrue(lengths, "No output collected") + self.assertEqual( + len(set(lengths)), + 1, + f"Output lengths vary across offsets: {lengths}", + ) + + @unittest.skipUnless(os.name == "posix", "Unix/Linux/macOS only") + def test_tail(self): + """Test multibyte decoding in Tail across buffer boundaries.""" + tail_lines = 3 + lengths = [] + output_buffer = [] + for offset in range(self.MAX_OFFSET): + output_buffer = [] + terminated = False + + def on_output(text): + nonlocal output_buffer + output_buffer.append(text) + + def on_terminate(_status): + nonlocal terminated + terminated = True + + cmd = self._multibyte_write_cmd(offset, count=tail_lines) + tail = client.Tail( + f'{sys.executable} -c "{cmd}"', + output_func=on_output, + termination_func=on_terminate, + ) + for _ in range(1000): + if terminated: + break + time.sleep(0.01) + tail.close() + for line in output_buffer: + if line.startswith("(Process terminated "): + continue + stripped = line.lstrip() + self.assertTrue( + stripped.startswith(self.TEXT), + f"offset {offset}: unexpected start: {stripped[:20]!r}", + ) + lengths.append(len(stripped)) + self.assertTrue(lengths, "No output collected") + self.assertEqual( + len(set(lengths)), + 1, + f"Output lengths vary across offsets: {lengths}", + ) + + if __name__ == "__main__": unittest.main()