Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 8 additions & 8 deletions py4j-python/src/py4j/clientserver.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)


Expand Down Expand Up @@ -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.
Expand All @@ -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)
Expand Down Expand Up @@ -588,15 +588,15 @@ 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(
command, self.stream, self.socket,
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))
Expand Down Expand Up @@ -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 +\
Expand All @@ -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):
Expand Down
42 changes: 25 additions & 17 deletions py4j-python/src/py4j/java_gateway.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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:]
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -2475,15 +2483,15 @@ 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.
authenticated = do_client_auth(
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))
Expand Down Expand Up @@ -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 +\
Expand All @@ -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


Expand Down Expand Up @@ -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
Expand Down
30 changes: 26 additions & 4 deletions py4j-python/src/py4j/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -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("\\", "\\\\").\
Expand Down Expand Up @@ -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":
Expand All @@ -234,8 +250,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):
Expand Down
40 changes: 40 additions & 0 deletions py4j-python/src/py4j/tests/protocol_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading