diff --git a/git_command.py b/git_command.py index 7add44747..abaf6a565 100644 --- a/git_command.py +++ b/git_command.py @@ -283,6 +283,7 @@ class GitCommand: bare=False, input=None, capture_stdout=False, + capture_stdout_bytes: bool = False, capture_stderr=False, merge_output=False, disable_editor=False, @@ -304,6 +305,12 @@ class GitCommand: self.cmdv = cmdv self.verify_command = verify_command self.stdout, self.stderr = None, None + if capture_stdout_bytes: + if merge_output: + raise ValueError( + "capture_stdout_bytes cannot be combined with merge_output" + ) + capture_stdout = True # Git on Windows wants its paths only using / for reliability. if platform_utils.isWindows(): @@ -347,6 +354,7 @@ class GitCommand: command, env, capture_stdout=capture_stdout, + capture_stdout_bytes=capture_stdout_bytes, capture_stderr=capture_stderr, merge_output=merge_output, ssh_proxy=ssh_proxy, @@ -380,6 +388,7 @@ class GitCommand: command, env, capture_stdout=False, + capture_stdout_bytes: bool = False, capture_stderr=False, merge_output=False, ssh_proxy=None, @@ -412,6 +421,10 @@ class GitCommand: # See go/tee-repo-stderr for more context. tee_stderr = False kwargs = {"encoding": "utf-8", "errors": "backslashreplace"} + if capture_stdout_bytes: + kwargs = {} + if isinstance(input, str): + input = input.encode("utf-8", "surrogateescape") if not (stdin or stdout or stderr): tee_stderr = True # stderr will be written back to sys.stderr even though it is @@ -490,6 +503,10 @@ class GitCommand: self.stderr = self._Tee(p.stderr, sys.stderr) else: self.stdout, self.stderr = p.communicate(input=input) + if capture_stdout_bytes and isinstance(self.stderr, bytes): + self.stderr = self.stderr.decode( + "utf-8", "backslashreplace" + ).replace("\r\n", "\n") finally: if ssh_proxy: ssh_proxy.remove_client(p) @@ -541,17 +558,35 @@ class GitCommand: env.pop(key, None) return env - def VerifyCommand(self): + def VerifyCommand(self) -> None: if self.rc == 0: return None - stdout = ( - "\n".join(self.stdout.split("\n")[:GIT_ERROR_STDOUT_LINES]) - if self.stdout - else None - ) + raw_stdout = self.stdout + if isinstance(raw_stdout, bytes): + first_records = re.split( + rb"\r\n|[\r\n\0]", raw_stdout, maxsplit=GIT_ERROR_STDOUT_LINES + )[:GIT_ERROR_STDOUT_LINES] + stdout = ( + "\n".join( + r.decode("utf-8", "backslashreplace") for r in first_records + ) + if raw_stdout + else None + ) + elif raw_stdout: + first_records = re.split( + r"\r\n|[\r\n\0]", raw_stdout, maxsplit=GIT_ERROR_STDOUT_LINES + )[:GIT_ERROR_STDOUT_LINES] + stdout = "\n".join(first_records) + else: + stdout = None + + raw_stderr = self.stderr + if isinstance(raw_stderr, bytes): + raw_stderr = raw_stderr.decode("utf-8", "backslashreplace") stderr = ( - "\n".join(self.stderr.split("\n")[:GIT_ERROR_STDERR_LINES]) - if self.stderr + "\n".join(raw_stderr.split("\n")[:GIT_ERROR_STDERR_LINES]) + if raw_stderr else None ) project = self.project.name if self.project else None diff --git a/tests/test_git_command.py b/tests/test_git_command.py index 2d8b0af61..3708d116e 100644 --- a/tests/test_git_command.py +++ b/tests/test_git_command.py @@ -123,6 +123,7 @@ class GitCommandStreamLogsTest(unittest.TestCase): """Tests the GitCommand class stderr log streaming cases.""" def setUp(self): + _ = git_command.user_agent.git self.mock_process = mock.MagicMock() self.mock_process.communicate.return_value = (None, None) self.mock_process.wait.return_value = 0 @@ -228,6 +229,141 @@ class GitCommandStreamLogsTest(unittest.TestCase): self.assertEqual(cmd.stderr, logs) +class GitCommandCaptureBytesTest(unittest.TestCase): + """Tests the GitCommand class byte capture cases.""" + + def setUp(self) -> None: + _ = git_command.user_agent.git + self.mock_process = mock.MagicMock() + self.mock_process.communicate.return_value = (None, None) + self.mock_process.wait.return_value = 0 + + self.mock_popen = mock.MagicMock() + self.mock_popen.return_value = self.mock_process + mock.patch("subprocess.Popen", self.mock_popen).start() + + def tearDown(self) -> None: + mock.patch.stopall() + + def test_captures_stdout_as_bytes(self) -> None: + self.mock_process.communicate.return_value = (b"\xff\x00", b"error\r\n") + + cmd = git_command.GitCommand( + None, + ["status"], + capture_stdout=True, + capture_stdout_bytes=True, + capture_stderr=True, + ) + + self.mock_popen.assert_called_once_with( + ["git", "status"], + cwd=None, + env=mock.ANY, + stdin=None, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + self.assertEqual(cmd.stdout, b"\xff\x00") + self.assertEqual(cmd.stderr, "error\n") + + def test_capture_stdout_bytes_auto_enables_capture_stdout(self) -> None: + self.mock_process.communicate.return_value = (b"output", b"") + + cmd = git_command.GitCommand( + None, + ["status"], + capture_stdout_bytes=True, + ) + + self.mock_popen.assert_called_once_with( + ["git", "status"], + cwd=None, + env=mock.ANY, + stdin=None, + stdout=subprocess.PIPE, + stderr=None, + ) + self.assertEqual(cmd.stdout, b"output") + + def test_capture_stdout_bytes_with_merge_output_raises(self) -> None: + with self.assertRaises(ValueError): + git_command.GitCommand( + None, + ["status"], + capture_stdout_bytes=True, + merge_output=True, + ) + + def test_captures_stdout_as_bytes_encodes_str_input(self) -> None: + self.mock_process.communicate.return_value = (b"output", b"") + + git_command.GitCommand( + None, + ["status"], + input="hello world", + capture_stdout_bytes=True, + ) + + self.mock_process.communicate.assert_called_once_with( + input=b"hello world" + ) + + def test_captures_stdout_as_bytes_encodes_surrogate_input(self) -> None: + self.mock_process.communicate.return_value = (b"output", b"") + + git_command.GitCommand( + None, + ["status"], + input="file_\udcff.txt", + capture_stdout_bytes=True, + ) + + self.mock_process.communicate.assert_called_once_with( + input=b"file_\xff.txt" + ) + + def test_captures_stdout_as_bytes_passes_bytes_input(self) -> None: + self.mock_process.communicate.return_value = (b"output", b"") + + git_command.GitCommand( + None, + ["status"], + input=b"raw_\xff.txt", + capture_stdout_bytes=True, + ) + + self.mock_process.communicate.assert_called_once_with( + input=b"raw_\xff.txt" + ) + + def test_verify_command_truncates_nul_delimited_stdout(self) -> None: + cmd = git_command.GitCommand( + None, + ["status"], + capture_stdout_bytes=True, + ) + cmd.rc = 1 + cmd.stdout = b"first_file\0second_file\0third_file" + cmd.stderr = "stderr" + with self.assertRaises(git_command.GitCommandError) as cm: + cmd.VerifyCommand() + self.assertEqual(cm.exception.git_stdout, "first_file") + + def test_verify_command_decodes_bytes_stdout(self) -> None: + cmd = git_command.GitCommand( + None, + ["status"], + capture_stdout_bytes=True, + ) + cmd.rc = 1 + cmd.stdout = b"error\xff\nline2" + cmd.stderr = "stderr" + with self.assertRaises(git_command.GitCommandError) as cm: + cmd.VerifyCommand() + self.assertEqual(cm.exception.git_stdout, "error\\xff") + + class GitCallUnitTest(unittest.TestCase): """Tests the _GitCall class (via git_command.git)."""