Skip to content

Commit 9461bea

Browse files
[3.15] gh-158803: Fix crash in bytes.join() on a concurrently mutated list (#159030)
[3.15] gh-158803: Fix crash in bytes.join() on a concurrently mutated list (GH-158910) In the free-threaded build, bytes.join() and bytearray.join() read items from the list with borrowed references and without holding its lock, so another thread could replace and free an item before it was increfed. Run the join under Py_BEGIN_CRITICAL_SECTION_SEQUENCE_FAST, as PyUnicode_Join() already does. (cherry picked from commit 05d80cc)
1 parent 626fb76 commit 9461bea

4 files changed

Lines changed: 58 additions & 13 deletions

File tree

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,33 @@
1+
import unittest
2+
from threading import Event
3+
from test.support import threading_helper
4+
5+
threading_helper.requires_working_threading(module=True)
6+
7+
8+
class BytesThreading(unittest.TestCase):
9+
def test_racing_join_replace(self):
10+
# gh-158803: join() must not use a list item that another thread
11+
# replaces (and frees) concurrently.
12+
lst = [bytes(10) for _ in range(100)]
13+
done = Event()
14+
15+
def writer():
16+
try:
17+
for _ in range(100):
18+
for i in range(len(lst)):
19+
lst[i] = bytearray(10) if i % 2 else bytes(10)
20+
finally:
21+
done.set()
22+
23+
def reader():
24+
while not done.is_set():
25+
b''.join(lst)
26+
b'-'.join(lst)
27+
bytearray().join(lst)
28+
29+
threading_helper.run_concurrently([writer] + [reader] * 4)
30+
31+
32+
if __name__ == "__main__":
33+
unittest.main()
Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
Fix a crash in :meth:`bytes.join` and :meth:`bytearray.join` in the
2+
:term:`free-threaded build` when another thread concurrently mutates the
3+
list being joined. Patch by Christian Aurich Zanettini Martins.

‎Objects/bytesobject.c‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
#include "pycore_bytesobject.h" // _PyBytes_Find(), _PyBytes_Repeat()
77
#include "pycore_call.h" // _PyObject_CallNoArgs()
88
#include "pycore_ceval.h" // _PyEval_GetBuiltin()
9+
#include "pycore_critical_section.h" // Py_BEGIN_CRITICAL_SECTION_SEQUENCE_FAST()
910
#include "pycore_format.h" // F_LJUST
1011
#include "pycore_freelist.h" // _Py_FREELIST_FREE()
1112
#include "pycore_global_objects.h"// _Py_GET_GLOBAL_OBJECT()

‎Objects/stringlib/join.h‎

Lines changed: 21 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
#endif
66

77
Py_LOCAL_INLINE(PyObject *)
8-
STRINGLIB(bytes_join)(PyObject *sep, PyObject *iterable)
8+
STRINGLIB(bytes_join_lock_held)(PyObject *sep, PyObject *seq)
99
{
1010
const char *sepstr = STRINGLIB_STR(sep);
1111
Py_ssize_t seplen = STRINGLIB_LEN(sep);
@@ -14,38 +14,29 @@ STRINGLIB(bytes_join)(PyObject *sep, PyObject *iterable)
1414
Py_ssize_t seqlen = 0;
1515
Py_ssize_t sz = 0;
1616
Py_ssize_t i, nbufs;
17-
PyObject *seq, *item;
17+
PyObject *item;
1818
Py_buffer *buffers = NULL;
1919
#define NB_STATIC_BUFFERS 10
2020
Py_buffer static_buffers[NB_STATIC_BUFFERS];
2121
#define GIL_THRESHOLD 1048576
2222
int drop_gil = 1;
2323
PyThreadState *save = NULL;
2424

25-
seq = PySequence_Fast(iterable, "can only join an iterable");
26-
if (seq == NULL) {
27-
return NULL;
28-
}
29-
3025
seqlen = PySequence_Fast_GET_SIZE(seq);
3126
if (seqlen == 0) {
32-
Py_DECREF(seq);
3327
return STRINGLIB_NEW(NULL, 0);
3428
}
3529
#if !STRINGLIB_MUTABLE
3630
if (seqlen == 1) {
3731
item = PySequence_Fast_GET_ITEM(seq, 0);
3832
if (STRINGLIB_CHECK_EXACT(item)) {
39-
Py_INCREF(item);
40-
Py_DECREF(seq);
41-
return item;
33+
return Py_NewRef(item);
4234
}
4335
}
4436
#endif
4537
if (seqlen > NB_STATIC_BUFFERS) {
4638
buffers = PyMem_NEW(Py_buffer, seqlen);
4739
if (buffers == NULL) {
48-
Py_DECREF(seq);
4940
PyErr_NoMemory();
5041
return NULL;
5142
}
@@ -155,13 +146,30 @@ STRINGLIB(bytes_join)(PyObject *sep, PyObject *iterable)
155146
error:
156147
res = NULL;
157148
done:
158-
Py_DECREF(seq);
159149
for (i = 0; i < nbufs; i++)
160150
PyBuffer_Release(&buffers[i]);
161151
if (buffers != static_buffers)
162152
PyMem_Free(buffers);
163153
return res;
164154
}
165155

156+
Py_LOCAL_INLINE(PyObject *)
157+
STRINGLIB(bytes_join)(PyObject *sep, PyObject *iterable)
158+
{
159+
PyObject *seq, *res;
160+
161+
seq = PySequence_Fast(iterable, "can only join an iterable");
162+
if (seq == NULL) {
163+
return NULL;
164+
}
165+
166+
Py_BEGIN_CRITICAL_SECTION_SEQUENCE_FAST(iterable);
167+
res = STRINGLIB(bytes_join_lock_held)(sep, seq);
168+
Py_END_CRITICAL_SECTION_SEQUENCE_FAST();
169+
170+
Py_DECREF(seq);
171+
return res;
172+
}
173+
166174
#undef NB_STATIC_BUFFERS
167175
#undef GIL_THRESHOLD

0 commit comments

Comments
 (0)