Skip to content

Commit 96dc88c

Browse files
committed
speedups: validate mask length
The lack of this check permitted a read of up to 3 bytes past the end of the string in some cases.
1 parent ff808b3 commit 96dc88c

3 files changed

Lines changed: 50 additions & 28 deletions

File tree

‎tornado/speedups.c‎

Lines changed: 40 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -2,63 +2,76 @@
22
#include <Python.h>
33
#include <stdint.h>
44

5-
static PyObject* websocket_mask(PyObject* self, PyObject* args) {
6-
const char* mask;
5+
static PyObject *websocket_mask(PyObject *self, PyObject *args)
6+
{
7+
const char *mask;
78
Py_ssize_t mask_len;
89
uint32_t uint32_mask;
910
uint64_t uint64_mask;
10-
const char* data;
11+
const char *data;
1112
Py_ssize_t data_len;
1213
Py_ssize_t i;
13-
PyObject* result;
14-
char* buf;
14+
PyObject *result;
15+
char *buf;
1516

16-
if (!PyArg_ParseTuple(args, "s#s#", &mask, &mask_len, &data, &data_len)) {
17+
if (!PyArg_ParseTuple(args, "s#s#", &mask, &mask_len, &data, &data_len))
18+
{
1719
return NULL;
1820
}
1921

20-
uint32_mask = ((uint32_t*)mask)[0];
22+
if (mask_len != 4)
23+
{
24+
PyErr_SetString(PyExc_ValueError, "mask must be 4 bytes");
25+
return NULL;
26+
}
27+
28+
uint32_mask = ((uint32_t *)mask)[0];
2129

2230
result = PyBytes_FromStringAndSize(NULL, data_len);
23-
if (!result) {
31+
if (!result)
32+
{
2433
return NULL;
2534
}
2635
buf = PyBytes_AsString(result);
2736

28-
if (sizeof(size_t) >= 8) {
37+
if (sizeof(size_t) >= 8)
38+
{
2939
uint64_mask = uint32_mask;
3040
uint64_mask = (uint64_mask << 32) | uint32_mask;
3141

32-
while (data_len >= 8) {
33-
((uint64_t*)buf)[0] = ((uint64_t*)data)[0] ^ uint64_mask;
42+
while (data_len >= 8)
43+
{
44+
((uint64_t *)buf)[0] = ((uint64_t *)data)[0] ^ uint64_mask;
3445
data += 8;
3546
buf += 8;
3647
data_len -= 8;
3748
}
3849
}
3950

40-
while (data_len >= 4) {
41-
((uint32_t*)buf)[0] = ((uint32_t*)data)[0] ^ uint32_mask;
51+
while (data_len >= 4)
52+
{
53+
((uint32_t *)buf)[0] = ((uint32_t *)data)[0] ^ uint32_mask;
4254
data += 4;
4355
buf += 4;
4456
data_len -= 4;
4557
}
4658

47-
for (i = 0; i < data_len; i++) {
59+
for (i = 0; i < data_len; i++)
60+
{
4861
buf[i] = data[i] ^ mask[i];
4962
}
5063

5164
return result;
5265
}
5366

54-
static int speedups_exec(PyObject *module) {
67+
static int speedups_exec(PyObject *module)
68+
{
5569
return 0;
5670
}
5771

5872
static PyMethodDef methods[] = {
59-
{"websocket_mask", websocket_mask, METH_VARARGS, ""},
60-
{NULL, NULL, 0, NULL}
61-
};
73+
{"websocket_mask", websocket_mask, METH_VARARGS, ""},
74+
{NULL, NULL, 0, NULL}};
6275

6376
static PyModuleDef_Slot slots[] = {
6477
{Py_mod_exec, speedups_exec},
@@ -68,19 +81,19 @@ static PyModuleDef_Slot slots[] = {
6881
#if (!defined(Py_LIMITED_API) && PY_VERSION_HEX >= 0x030d0000) || Py_LIMITED_API >= 0x030d0000
6982
{Py_mod_gil, Py_MOD_GIL_NOT_USED},
7083
#endif
71-
{0, NULL}
72-
};
84+
{0, NULL}};
7385

7486
static struct PyModuleDef speedupsmodule = {
75-
PyModuleDef_HEAD_INIT,
76-
"speedups",
77-
NULL,
78-
0,
79-
methods,
80-
slots,
87+
PyModuleDef_HEAD_INIT,
88+
"speedups",
89+
NULL,
90+
0,
91+
methods,
92+
slots,
8193
};
8294

8395
PyMODINIT_FUNC
84-
PyInit_speedups(void) {
96+
PyInit_speedups(void)
97+
{
8598
return PyModuleDef_Init(&speedupsmodule);
8699
}

‎tornado/test/websocket_test.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -794,6 +794,13 @@ def test_mask(self: typing.Any):
794794
b"\xff\xfa\xff\xff\xfb\xfe",
795795
)
796796

797+
def test_length_validation(self: typing.Any):
798+
# Test all lengths of mask that are not 4 bytes.
799+
for mask in (b"", b"a", b"ab", b"abc", b"abcde", b"abcdef"):
800+
with self.subTest(mask=mask):
801+
with self.assertRaises(ValueError):
802+
self.mask(mask, b"data asdf")
803+
797804

798805
class PythonMaskFunctionTest(MaskFunctionMixin):
799806
def mask(self, mask, data):

‎tornado/util.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -145,7 +145,7 @@ def exec_in(
145145

146146

147147
def raise_exc_info(
148-
exc_info: Tuple[Optional[type], Optional[BaseException], Optional["TracebackType"]]
148+
exc_info: Tuple[Optional[type], Optional[BaseException], Optional["TracebackType"]],
149149
) -> typing.NoReturn:
150150
try:
151151
if exc_info[1] is not None:
@@ -418,6 +418,8 @@ def _websocket_mask_python(mask: bytes, data: bytes) -> bytes:
418418
419419
This pure-python implementation may be replaced by an optimized version when available.
420420
"""
421+
if len(mask) != 4:
422+
raise ValueError("mask must be 4 bytes")
421423
mask_arr = array.array("B", mask)
422424
unmasked_arr = array.array("B", data)
423425
for i in range(len(data)):

0 commit comments

Comments
 (0)