diff --git a/Lib/test/test_bytes.py b/Lib/test/test_bytes.py index 419bee5583de47..ee18e2de2403a2 100644 --- a/Lib/test/test_bytes.py +++ b/Lib/test/test_bytes.py @@ -1138,6 +1138,12 @@ def test_translate(self): c = b.translate(None, b'e') self.assertEqual(c, b'hllo') + # short inputs starting with NUL bytes + table = bytes.maketrans(b'\x00', b'Z') + for data in b'\x00', b'\x00' * 8, b'\x00' * 8 + b'a' * 247: + c = self.type2test(data).translate(table) + self.assertEqual(c, data.replace(b'\x00', b'Z')) + # test delete as a keyword argument c = b.translate(rosetta, delete=b'') self.assertEqual(c, b'helle') diff --git a/Objects/bytesobject.c b/Objects/bytesobject.c index 91cbfa23e30b20..99489f876fc33e 100644 --- a/Objects/bytesobject.c +++ b/Objects/bytesobject.c @@ -2274,7 +2274,7 @@ bytes_translate_impl(PyBytesObject *self, PyObject *table, const char *table_chars; Py_ssize_t i, c, changed = 0; PyObject *input_obj = (PyObject*)self; - const char *output_start, *del_table_chars=NULL; + const char *del_table_chars=NULL; Py_ssize_t inlen, tablen, dellen = 0; int trans_table[256]; @@ -2326,7 +2326,7 @@ bytes_translate_impl(PyBytesObject *self, PyObject *table, PyBuffer_Release(&table_view); return NULL; } - output_start = output = PyBytesWriter_GetData(writer); + output = PyBytesWriter_GetData(writer); input = PyBytes_AS_STRING(input_obj); if (dellen == 0 && table_chars != NULL) { @@ -2341,7 +2341,8 @@ bytes_translate_impl(PyBytesObject *self, PyObject *table, /* We save this check until the end so that the compiler will */ /* unroll the loop above leading to MUCH faster code. */ if (result != NULL && PyBytes_CheckExact(input_obj)) { - if (memcmp(PyBytes_AS_STRING(input_obj), output_start, inlen) == 0) { + if (memcmp(PyBytes_AS_STRING(input_obj), + PyBytes_AS_STRING(result), inlen) == 0) { Py_SETREF(result, Py_NewRef(input_obj)); } }