From 62361831be4a9580f03ed5f413e1799f2c9199ed Mon Sep 17 00:00:00 2001 From: Ruslan Dautkhanov Date: Wed, 20 May 2026 23:30:27 -0600 Subject: [PATCH 1/2] perf: drop the byte-list intermediate in decode_bytearray (closes #570) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Per @PaperTsar's analysis in issue #570, the list comprehension [b for b in standard_b64decode(...)] allocates one PyObject per byte and then reconstructs bytes from that list — pure overhead now that Python 2 is no longer a target. standard_b64decode returns bytes directly; the bytes() wrapper preserves the contract. ~7.5x speedup on 256KB payloads (microbench); CodSpeed run in this PR gives the authoritative scenario-level number. Closes #570. Co-authored-by: Isaac --- py4j-python/src/py4j/protocol.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/py4j-python/src/py4j/protocol.py b/py4j-python/src/py4j/protocol.py index 059d926d..38599701 100644 --- a/py4j-python/src/py4j/protocol.py +++ b/py4j-python/src/py4j/protocol.py @@ -234,8 +234,14 @@ def encode_bytearray(barray): def decode_bytearray(encoded): - new_bytes = bytes(encoded, encoding="ascii") - return bytes([b for b in standard_b64decode(new_bytes)]) + # Per @PaperTsar's analysis in issue #570: the prior + # implementation built a Python list of ints (one PyObject per + # byte) then reconstructed bytes from that list — pure overhead + # now that Python 2 is no longer a target. standard_b64decode + # already returns bytes; the bytes() wrapper preserves the return- + # type contract while skipping the intermediate list. ~7.5x on + # 256KB payloads in microbenchmarks. + return bytes(standard_b64decode(encoded.encode("ascii"))) def is_python_proxy(parameter): From e7b1485eed81df27214fae5cfc747c7c37899a77 Mon Sep 17 00:00:00 2001 From: Ruslan Dautkhanov Date: Thu, 21 May 2026 10:55:43 -0600 Subject: [PATCH 2/2] perf: drop smart_decode dispatch on hot recv + minor protocol cleanups MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Builds on the decode_bytearray fix in the previous commit by adopting the hot-path perf shortcuts from @markjm's PR py4j/py4j#575 — but carefully, avoiding the bytes/str regression that caused testGatewayAuth to fail on that PR (see review at py4j/py4j#575 (comment) 4484469752). ## What's adopted from #575 * **socket recv decoded at the source** — replace ``smart_decode(self.stream.readline()[:-1])`` with ``self.stream.readline()[:-1].decode("utf-8")`` at every hot read site: * ``java_gateway.py``: ``GatewayConnection.send_command`` (every JavaGateway round-trip), ``do_client_auth`` (auth handshake), 4 sites in the callback-server ``run`` / ``_call_proxy`` / ``_get_params`` paths, and ``OutputConsumer.run`` (stdout consumer thread). * ``clientserver.py``: 7 sites in ClientServer send/recv + proxy command + ``_get_params``. Each saves the smart_decode dispatch (one isinstance call per read); on the per-call path this compounds quickly. * **encode_float**: drop ``smart_decode(repr(float))`` — ``str(float)`` returns the same shortest-roundtrip repr on Python 3, no dispatch. * **str() in cold sites**: replace ``smart_decode(addr/port/id)`` with ``str(...)`` for finalizer-key construction, logging of getsockname(), and python-proxy-id generation. These were never bytes inputs; smart_decode was always doing the str() fallback. * **launch_gateway** auth-token: decode at the source (``proc.stdout.readline()[:-len(os.linesep)].decode("utf-8")``). Without this, the rest of the auth flow compares bytes to str and silently fails authentication. ## What's deliberately NOT adopted from #575 * **smart_decode in ``escape_new_line`` is KEPT.** Dropping it (as #575 proposed) makes ``bytes.replace("str", "str")`` raise TypeError. This was the testGatewayAuth regression I flagged on py4j/py4j#575. The replace chain is fast enough that the smart_decode dispatch isn't a hot-path concern, and the safety net is load-bearing for any caller that hands escape_new_line bytes (auth tokens, legacy callers). Docstring updated to make the rationale explicit so a future refactor doesn't accidentally drop it. * **decode_bytearray return type stays ``bytes`` (not ``bytearray``).** #575 changed it to ``bytearray``, but the original (and current) return contract was ``bytes`` (via the ``bytearray2 = bytes`` alias on Python 3). A return-type change risks breaking downstream consumers (e.g. PySpark's binary deserialization). The decode_bytearray fix from issue #570 is already in place from the previous commit; this commit doesn't touch it. ## Tests added * **EscapeNewLineBytesInputSafetyTest** (4 tests) — pin the bytes-input contract on escape_new_line. ASCII bytes / str / UTF-8 bytes / empty bytes all flow through cleanly. If a future refactor drops the smart_decode safety net, these tests fail immediately — catching the regression before the JVM-bound testGatewayAuth integration test would. ## Existing integration tests that validate this change * ``GatewayLauncherTest.testGatewayAuth`` (java_gateway_test.py) — exercises ``launch_gateway(enable_auth=True)`` end-to-end. The bytes-stream auth-token read passes through the new ``.decode("utf-8")`` path. * ``PythonEntryPointTest.test_python_entry_point_with_auth`` (java_callback_test.py) — exercises the callback-with-auth path through all 4 modified read sites in the callback server. Both pass on the full Python x Java x OS matrix. Credit to @markjm for the perf shortcuts in #575. Co-authored-by: Isaac --- py4j-python/src/py4j/clientserver.py | 16 ++++---- py4j-python/src/py4j/java_gateway.py | 42 ++++++++++++--------- py4j-python/src/py4j/protocol.py | 20 +++++++++- py4j-python/src/py4j/tests/protocol_test.py | 40 ++++++++++++++++++++ 4 files changed, 91 insertions(+), 27 deletions(-) diff --git a/py4j-python/src/py4j/clientserver.py b/py4j-python/src/py4j/clientserver.py index 1517a523..88a31c21 100644 --- a/py4j-python/src/py4j/clientserver.py +++ b/py4j-python/src/py4j/clientserver.py @@ -23,7 +23,7 @@ disable_nagle) from py4j import protocol as proto from py4j.protocol import ( - Py4JError, Py4JNetworkError, smart_decode, get_command_part, + Py4JError, Py4JNetworkError, get_command_part, get_return_value, Py4JAuthenticationError) @@ -530,7 +530,7 @@ def send_command(self, command): try: while True: - answer = smart_decode(self.stream.readline()[:-1]) + answer = self.stream.readline()[:-1].decode("utf-8") logger.debug("Answer received: {0}".format(answer)) # Happens when a the other end is dead. There might be an empty # answer before the socket raises an error. @@ -541,7 +541,7 @@ def send_command(self, command): return answer[1:] else: command = answer - obj_id = smart_decode(self.stream.readline())[:-1] + obj_id = self.stream.readline()[:-1].decode("utf-8") if command == proto.CALL_PROXY_COMMAND_NAME: return_message = self._call_proxy(obj_id, self.stream) @@ -588,7 +588,7 @@ def wait_for_commands(self): authenticated = self.python_parameters.auth_token is None try: while True: - command = smart_decode(self.stream.readline())[:-1] + command = self.stream.readline()[:-1].decode("utf-8") if not authenticated: # Will raise an exception if auth fails in any way. authenticated = do_client_auth( @@ -596,7 +596,7 @@ def wait_for_commands(self): self.python_parameters.auth_token) continue - obj_id = smart_decode(self.stream.readline())[:-1] + obj_id = self.stream.readline()[:-1].decode("utf-8") logger.info( "Received command {0} on object id {1}". format(command, obj_id)) @@ -637,7 +637,7 @@ def _call_proxy(self, obj_id, input): get_command_part('Object ID unknown', self.pool) try: - method = smart_decode(input.readline())[:-1] + method = input.readline()[:-1].decode("utf-8") params = self._get_params(input) return_value = getattr(self.pool[obj_id], method)(*params) return proto.RETURN_MESSAGE + proto.SUCCESS +\ @@ -657,11 +657,11 @@ def _call_proxy(self, obj_id, input): def _get_params(self, input): params = [] - temp = smart_decode(input.readline())[:-1] + temp = input.readline()[:-1].decode("utf-8") while temp != proto.END: param = get_return_value("y" + temp, self.java_client) params.append(param) - temp = smart_decode(input.readline())[:-1] + temp = input.readline()[:-1].decode("utf-8") return params def __del__(self): diff --git a/py4j-python/src/py4j/java_gateway.py b/py4j-python/src/py4j/java_gateway.py index 3b5a3d2f..0792fa22 100644 --- a/py4j-python/src/py4j/java_gateway.py +++ b/py4j-python/src/py4j/java_gateway.py @@ -31,7 +31,7 @@ Py4JError, Py4JJavaError, Py4JNetworkError, Py4JAuthenticationError, get_command_part, get_return_value, - register_output_converter, smart_decode, escape_new_line, + register_output_converter, escape_new_line, is_fatal_error, is_error, unescape_new_line, get_error_message, compute_exception_message) from py4j.signals import Signal @@ -361,10 +361,13 @@ def launch_gateway(port=0, jarpath="", classpath="", javaopts=[], # ephemeral ports) _port = int(proc.stdout.readline()) - # Read the auth token from the server if enabled. + # Read the auth token from the server if enabled. stdout is in + # binary mode by default; decode here so the rest of the auth flow + # (which uses string equality, escape_new_line, etc.) sees str. _auth_token = None if enable_auth: - _auth_token = proc.stdout.readline()[:-len(os.linesep)] + _auth_token = proc.stdout.readline()[:-len(os.linesep)].decode( + "utf-8") # Start consumer threads so process does not deadlock/hangs OutputConsumer( @@ -638,7 +641,7 @@ def do_client_auth(command, input_stream, sock, auth_token): raise Py4JAuthenticationError("Expected {}, received {}.".format( proto.AUTH_COMMAND_NAME, command)) - client_token = smart_decode(input_stream.readline()[:-1]) + client_token = input_stream.readline()[:-1].decode("utf-8") # Remove the END marker input_stream.readline() if auth_token == client_token: @@ -669,8 +672,8 @@ def _garbage_collect_object(gateway_client, target_id): try: try: ThreadSafeFinalizer.remove_finalizer( - smart_decode(gateway_client.address) + - smart_decode(gateway_client.port) + + str(gateway_client.address) + + str(gateway_client.port) + target_id) gateway_client.garbage_collect_object(target_id) except Exception: @@ -747,7 +750,8 @@ def _pipe_fd(self, line): def run(self): lines_iterator = iter(self.stream.readline, b"") for line in lines_iterator: - self.redirect_func(smart_decode(line)) + # The sentinel b"" above pins line to bytes; decode directly. + self.redirect_func(line.decode("utf-8")) class ProcessConsumer(Thread): @@ -1278,7 +1282,11 @@ def send_command(self, command): "Error while sending", e, proto.ERROR_ON_SEND) try: - answer = smart_decode(self.stream.readline()[:-1]) + # Stream is opened in binary mode (socket.makefile("rb")), + # so readline() returns bytes; decode at the source rather + # than dispatch through smart_decode's isinstance check. + # Every JavaGateway call hits this — the saving compounds. + answer = self.stream.readline()[:-1].decode("utf-8") logger.debug("Answer received: {0}".format(answer)) if answer.startswith(proto.RETURN_MESSAGE): answer = answer[1:] @@ -1427,8 +1435,8 @@ def __init__(self, target_id, gateway_client): self._fully_populated = False self._gateway_doc = None - key = smart_decode(self._gateway_client.address) +\ - smart_decode(self._gateway_client.port) +\ + key = str(self._gateway_client.address) +\ + str(self._gateway_client.port) +\ self._target_id if self._gateway_client.gateway_property.enable_memory_management: @@ -2352,7 +2360,7 @@ def run(self): self.server_socket.listen(5) logger.info( "Socket listening on {0}". - format(smart_decode(self.server_socket.getsockname()))) + format(self.server_socket.getsockname())) server_started.send( self, server=self) @@ -2475,7 +2483,7 @@ def run(self): authenticated = self.callback_server_parameters.auth_token is None try: while True: - command = smart_decode(self.input.readline())[:-1] + command = self.input.readline()[:-1].decode("utf-8") if not authenticated: token = self.callback_server_parameters.auth_token # Will raise an exception if auth fails in any way. @@ -2483,7 +2491,7 @@ def run(self): command, self.input, self.socket, token) continue - obj_id = smart_decode(self.input.readline())[:-1] + obj_id = self.input.readline()[:-1].decode("utf-8") logger.info( "Received command {0} on object id {1}". format(command, obj_id)) @@ -2540,7 +2548,7 @@ def _call_proxy(self, obj_id, input): get_command_part('Object ID unknown', self.pool) try: - method = smart_decode(input.readline())[:-1] + method = input.readline()[:-1].decode("utf-8") params = self._get_params(input) return_value = getattr(self.pool[obj_id], method)(*params) return proto.RETURN_MESSAGE + proto.SUCCESS +\ @@ -2559,11 +2567,11 @@ def _call_proxy(self, obj_id, input): def _get_params(self, input): params = [] - temp = smart_decode(input.readline())[:-1] + temp = input.readline()[:-1].decode("utf-8") while temp != proto.END: param = get_return_value("y" + temp, self.gateway_client) params.append(param) - temp = smart_decode(input.readline())[:-1] + temp = input.readline()[:-1].decode("utf-8") return params @@ -2596,7 +2604,7 @@ def put(self, object, force_id=None): if force_id: id = force_id else: - id = proto.PYTHON_PROXY_PREFIX + smart_decode(self.next_id) + id = proto.PYTHON_PROXY_PREFIX + str(self.next_id) self.next_id += 1 self.dict[id] = object return id diff --git a/py4j-python/src/py4j/protocol.py b/py4j-python/src/py4j/protocol.py index 38599701..bb14d36a 100644 --- a/py4j-python/src/py4j/protocol.py +++ b/py4j-python/src/py4j/protocol.py @@ -173,9 +173,22 @@ def escape_new_line(original): Backslashes are also escaped by another backslash. - :param original: the string to escape + :param original: the string to escape (str or bytes; bytes inputs + are decoded via smart_decode for backward compatibility — see + below). :rtype: an escaped string + + .. note:: + The internal ``smart_decode(original)`` is **load-bearing**: it + accepts bytes inputs that some legacy callers (and any code path + that forgot to decode at the socket boundary) might still + produce. Removing it makes ``bytes.replace("str", "str")`` + raise ``TypeError`` — see PR #575 review for the auth-token + regression this guarded against. The replacement chain is fast + enough that the smart_decode dispatch is not a hot-path concern; + all py4j-internal callers already pass str, so the type check + is a single isinstance hit. """ if original: return smart_decode(original).replace("\\", "\\\\").\ @@ -215,7 +228,10 @@ def smart_decode(s): def encode_float(float_value): - float_str = smart_decode(repr(float_value)) + # str(float) on Python 3 already returns the same shortest- + # roundtrip repr that smart_decode(repr(...)) was producing on + # py2; smart_decode here was a no-op dispatcher. + float_str = str(float_value) if float_str == "-inf": float_str = JAVA_NEGATIVE_INFINITY elif float_str == "inf": diff --git a/py4j-python/src/py4j/tests/protocol_test.py b/py4j-python/src/py4j/tests/protocol_test.py index 40f876f4..59f0ae7d 100644 --- a/py4j-python/src/py4j/tests/protocol_test.py +++ b/py4j-python/src/py4j/tests/protocol_test.py @@ -311,5 +311,45 @@ def test_only_special_chars(self): self.assertEqual(self._roundtrip(s), s) +class EscapeNewLineBytesInputSafetyTest(unittest.TestCase): + """Pins the bytes-input safety contract on escape_new_line. + + escape_new_line's ``smart_decode(original)`` is a defensive measure + that lets bytes inputs pass through cleanly — see the docstring of + escape_new_line for the full rationale. PR #575's perf review + proposed dropping this smart_decode; doing so makes + ``bytes.replace("str", "str")`` raise TypeError, breaking the + auth-token round-trip path (testGatewayAuth) and any other path + that hands escape_new_line a bytes input without explicit decoding. + + If a future refactor drops smart_decode from escape_new_line, these + tests fail immediately — catching the regression before CI's + integration tests need to spin up a JVM.""" + + def test_bytes_input_decoded_as_utf8(self): + # ASCII bytes round-trip through escape_new_line as if they + # were str — smart_decode does the conversion. + result = escape_new_line(b"hello\nworld") + self.assertEqual(result, "hello\\nworld") + + def test_str_input_passes_through(self): + # str inputs are the common case; smart_decode is a single + # isinstance hit for these. + result = escape_new_line("hello\nworld") + self.assertEqual(result, "hello\\nworld") + + def test_bytes_input_with_utf8_payload(self): + # Non-ASCII bytes (UTF-8 encoded) decode correctly via + # smart_decode("utf-8") — auth tokens or other identifiers + # may contain UTF-8 bytes if read from stdout in binary mode. + s = "\u4e2d\u6587" # "中文" + result = escape_new_line(s.encode("utf-8")) + self.assertEqual(result, s) + + def test_empty_bytes_passes_through(self): + # The falsy-passthrough branch handles both b"" and "". + self.assertEqual(escape_new_line(b""), b"") + + if __name__ == "__main__": unittest.main()