Skip to content

Commit 731dd23

Browse files
committed
[3.14] gh-155835: fix digest_size and block_size data races on SHA-3 objects (GH-155838)
(cherry picked from commit 90ac539) Co-authored-by: Bénédikt Tran <10796600+picnixz@users.noreply.github.com>
1 parent a9152fb commit 731dd23

3 files changed

Lines changed: 184 additions & 16 deletions

File tree

Lib/test/test_hashlib.py

Lines changed: 161 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,8 @@
1919
import threading
2020
import unittest
2121
import warnings
22+
from functools import partial
23+
from operator import attrgetter
2224
from test import support
2325
from test.support import _4G, bigmemtest
2426
from test.support import hashlib_helper
@@ -53,18 +55,39 @@
5355
def get_fips_mode():
5456
return 0
5557

58+
59+
try:
60+
import _md5
61+
except ImportError:
62+
_md5 = None
63+
requires_md5 = unittest.skipUnless(_md5, 'requires _md5')
64+
65+
5666
try:
5767
import _blake2
5868
except ImportError:
5969
_blake2 = None
60-
6170
requires_blake2 = unittest.skipUnless(_blake2, 'requires _blake2')
6271

72+
73+
try:
74+
import _sha1
75+
except ImportError:
76+
_sha1 = None
77+
requires_sha1 = unittest.skipUnless(_sha1, 'requires _sha1')
78+
79+
80+
try:
81+
import _sha2
82+
except ImportError:
83+
_sha2 = None
84+
requires_sha2 = unittest.skipUnless(_sha2, 'requires _sha2')
85+
86+
6387
try:
6488
import _sha3
6589
except ImportError:
6690
_sha3 = None
67-
6891
requires_sha3 = unittest.skipUnless(_sha3, 'requires _sha3')
6992

7093

@@ -1301,5 +1324,141 @@ def readable(self):
13011324
hashlib.file_digest(NonBlocking(), hashlib.sha256)
13021325

13031326

1327+
@unittest.skipUnless(hasattr(hashlib, 'scrypt'), 'requires OpenSSL 1.1+')
1328+
@unittest.skipIf(get_fips_mode(), reason="scrypt is blocked in FIPS mode")
1329+
class TestScrypt(unittest.TestCase):
1330+
1331+
scrypt_test_vectors = [
1332+
(b'', b'', 16, 1, 1, unhexlify('77d6576238657b203b19ca42c18a0497f16b4844e3074ae8dfdffa3fede21442fcd0069ded0948f8326a753a0fc81f17e8d3e0fb2e0d3628cf35e20c38d18906')),
1333+
(b'password', b'NaCl', 1024, 8, 16, unhexlify('fdbabe1c9d3472007856e7190d01e9fe7c6ad7cbc8237830e77376634b3731622eaf30d92e22a3886ff109279d9830dac727afb94a83ee6d8360cbdfa2cc0640')),
1334+
(b'pleaseletmein', b'SodiumChloride', 16384, 8, 1, unhexlify('7023bdcb3afd7348461c06cd81fd38ebfda8fbba904f8e3ea9b543f6545da1f2d5432955613f0fcf62d49705242a9af9e61e85dc0d651e40dfcf017b45575887')),
1335+
]
1336+
1337+
def test_scrypt(self):
1338+
for password, salt, n, r, p, expected in self.scrypt_test_vectors:
1339+
result = hashlib.scrypt(password, salt=salt, n=n, r=r, p=p)
1340+
self.assertEqual(result, expected)
1341+
1342+
# these parameters must be valid
1343+
hashlib.scrypt(b'password', salt=b'salt', n=2, r=8, p=1)
1344+
hashlib.scrypt(b'password', salt=b'salt', n=2, r=8, p=1, maxmem=0)
1345+
hashlib.scrypt(b'password', salt=b'salt', n=2, r=8, p=1, dklen=1)
1346+
1347+
def test_scrypt_types(self):
1348+
# password and salt must be bytes-like
1349+
with self.assertRaises(TypeError):
1350+
hashlib.scrypt('password', salt=b'salt', n=2, r=8, p=1)
1351+
with self.assertRaises(TypeError):
1352+
hashlib.scrypt(b'password', salt='salt', n=2, r=8, p=1)
1353+
# require keyword args
1354+
with self.assertRaises(TypeError):
1355+
hashlib.scrypt(b'password')
1356+
with self.assertRaises(TypeError):
1357+
hashlib.scrypt(b'password', b'salt')
1358+
with self.assertRaises(TypeError):
1359+
hashlib.scrypt(b'password', 2, 8, 1, salt=b'salt')
1360+
1361+
def test_scrypt_validate(self):
1362+
def scrypt(password=b"password", /, **kwargs):
1363+
# overwrite well-defined parameters with bad ones
1364+
kwargs = dict(salt=b'salt', n=2, r=8, p=1) | kwargs
1365+
return hashlib.scrypt(password, **kwargs)
1366+
1367+
for param_name in ('n', 'r', 'p', 'maxmem', 'dklen'):
1368+
param = {param_name: None}
1369+
with self.subTest(**param):
1370+
self.assertRaises(TypeError, scrypt, **param)
1371+
1372+
self.assertRaises(ValueError, scrypt, n=0)
1373+
self.assertRaises(ValueError, scrypt, n=-1)
1374+
self.assertRaises(ValueError, scrypt, n=1)
1375+
1376+
self.assertRaises(ValueError, scrypt, r=0)
1377+
self.assertRaises(ValueError, scrypt, r=-1)
1378+
1379+
self.assertRaises(ValueError, scrypt, p=-1)
1380+
self.assertRaises(ValueError, scrypt, p=0)
1381+
1382+
self.assertRaises(ValueError, scrypt, maxmem=-1)
1383+
# OpenSSL hard limit for 'maxmem' is an 'uint64_t' but for now,
1384+
# we do not use the 'uint64' Clinic converter but the 'long' one.
1385+
self.assertRaises(OverflowError, scrypt, maxmem=(1 << 64))
1386+
# Historically, Python allowed 'maxmem' to be at most INT_MAX,
1387+
# which is at most 2**32-1 (on Windows, sizeof(long) == 4, so
1388+
# an OverflowError will be raised instead of a ValueError).
1389+
numeric_exc_types = (OverflowError, ValueError)
1390+
self.assertRaises(numeric_exc_types, scrypt, maxmem=(1 << 32))
1391+
1392+
self.assertRaises(ValueError, scrypt, dklen=-1)
1393+
self.assertRaises(ValueError, scrypt, dklen=0)
1394+
MAX_DKLEN = ((1 << 32) - 1) * 32 # see RFC 7914
1395+
self.assertRaises(numeric_exc_types, scrypt, dklen=MAX_DKLEN + 1)
1396+
1397+
1398+
@threading_helper.requires_working_threading()
1399+
class TestTSAN(unittest.TestCase):
1400+
1401+
@threading_helper.reap_threads
1402+
def check_attribute(self, write, read, expected, nthreads=8):
1403+
ready = threading.Event()
1404+
barrier = threading.Barrier(nthreads)
1405+
1406+
def writer():
1407+
barrier.wait()
1408+
while not ready.is_set():
1409+
write()
1410+
1411+
def reader():
1412+
barrier.wait()
1413+
while not ready.is_set():
1414+
self.assertEqual(read(), expected)
1415+
1416+
targets = [writer if i % 2 else reader for i in range(nthreads)]
1417+
workers = [threading.Thread(target=target) for target in targets]
1418+
with threading_helper.start_threads(workers, unlock=ready.set):
1419+
pass
1420+
1421+
def check_HACL_attribute(self, module, version, attrname):
1422+
blob = b"A" * 65536
1423+
obj = getattr(module, version)()
1424+
update = partial(obj.update, blob)
1425+
read = attrgetter(attrname)
1426+
self.check_attribute(update, partial(read, obj), read(obj))
1427+
1428+
@requires_md5
1429+
@support.subTests("attrname", ["block_size", "digest_size"])
1430+
def test_HACL_md5_attributes(self, attrname):
1431+
self.check_HACL_attribute(_md5, "md5", attrname)
1432+
1433+
@requires_sha1
1434+
@support.subTests("attrname", ["block_size", "digest_size"])
1435+
def test_HACL_sha1_attributes(self, attrname):
1436+
self.check_HACL_attribute(_sha1, "sha1", attrname)
1437+
1438+
@requires_sha2
1439+
@support.subTests("size", [224, 256, 384, 512])
1440+
@support.subTests("attrname", ["block_size", "digest_size"])
1441+
def test_HACL_sha2_attributes(self, size, attrname):
1442+
self.check_HACL_attribute(_sha2, f"sha{size}", attrname)
1443+
1444+
@requires_sha3
1445+
@support.subTests("size", [224, 256, 384, 512])
1446+
@support.subTests(
1447+
"attrname",
1448+
["block_size", "digest_size", "_capacity_bits", "_rate_bits"],
1449+
)
1450+
def test_HACL_sha3_attributes(self, size, attrname):
1451+
self.check_HACL_attribute(_sha3, f"sha3_{size}", attrname)
1452+
1453+
@requires_sha3
1454+
@support.subTests("size", [128, 256])
1455+
@support.subTests(
1456+
"attrname",
1457+
["block_size", "digest_size", "_capacity_bits", "_rate_bits"],
1458+
)
1459+
def test_HACL_shake_attributes(self, size, attrname):
1460+
self.check_HACL_attribute(_sha3, f"shake_{size}", attrname)
1461+
1462+
13041463
if __name__ == "__main__":
13051464
unittest.main()
Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
:mod:`hashlib`: Fix data races when accessing
2+
:attr:`~hashlib.hash.digest_size` and :attr:`~hashlib.hash.block_size` on
3+
SHA-3 objects. Patch by Bénédikt Tran.

Modules/sha3module.c

Lines changed: 20 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,12 @@ typedef struct {
6464
bool use_mutex;
6565
PyMutex mutex;
6666
Hacl_Hash_SHA3_state_t *hash_state;
67+
// HACL* update functions entirely replace the state, which can lead
68+
// to races on the free-threaded build. Since the kind of hash is static,
69+
// we can store its corresponding metadata once.
70+
uint32_t digest_size;
71+
uint32_t block_size;
72+
int is_shake;
6773
} SHA3object;
6874

6975
#define _SHA3object_CAST(op) ((SHA3object *)(op))
@@ -78,7 +84,7 @@ newSHA3object(PyTypeObject *type)
7884
return NULL;
7985
}
8086
HASHLIB_INIT_MUTEX(newobj);
81-
87+
newobj->digest_size = newobj->block_size = 0;
8288
PyObject_GC_Track(newobj);
8389
return newobj;
8490
}
@@ -161,6 +167,11 @@ py_sha3_new_impl(PyTypeObject *type, PyObject *data_obj, int usedforsecurity,
161167
goto error;
162168
}
163169

170+
// set the metadata once we know that the state is valid
171+
int is_shake = Hacl_Hash_SHA3_is_shake(self->hash_state);
172+
self->digest_size = is_shake ? 0 : Hacl_Hash_SHA3_hash_len(self->hash_state);
173+
self->block_size = Hacl_Hash_SHA3_block_len(self->hash_state);
174+
164175
if (data) {
165176
GET_BUFFER_VIEW_OR_ERROR(data, &buf, goto error);
166177
if (buf.len >= HASHLIB_GIL_MINSIZE) {
@@ -245,6 +256,8 @@ _sha3_sha3_224_copy_impl(SHA3object *self)
245256
Py_DECREF(newobj);
246257
return PyErr_NoMemory();
247258
}
259+
newobj->digest_size = self->digest_size;
260+
newobj->block_size = self->block_size;
248261
return (PyObject *)newobj;
249262
}
250263

@@ -265,8 +278,7 @@ _sha3_sha3_224_digest_impl(SHA3object *self)
265278
ENTER_HASHLIB(self);
266279
(void)Hacl_Hash_SHA3_digest(self->hash_state, digest);
267280
LEAVE_HASHLIB(self);
268-
return PyBytes_FromStringAndSize((const char *)digest,
269-
Hacl_Hash_SHA3_hash_len(self->hash_state));
281+
return PyBytes_FromStringAndSize((const char *)digest, self->digest_size);
270282
}
271283

272284

@@ -284,8 +296,7 @@ _sha3_sha3_224_hexdigest_impl(SHA3object *self)
284296
ENTER_HASHLIB(self);
285297
(void)Hacl_Hash_SHA3_digest(self->hash_state, digest);
286298
LEAVE_HASHLIB(self);
287-
return _Py_strhex((const char *)digest,
288-
Hacl_Hash_SHA3_hash_len(self->hash_state));
299+
return _Py_strhex((const char *)digest, self->digest_size);
289300
}
290301

291302

@@ -337,8 +348,7 @@ static PyObject *
337348
SHA3_get_block_size(PyObject *op, void *Py_UNUSED(closure))
338349
{
339350
SHA3object *self = _SHA3object_CAST(op);
340-
uint32_t rate = Hacl_Hash_SHA3_block_len(self->hash_state);
341-
return PyLong_FromLong(rate);
351+
return PyLong_FromLong(self->block_size);
342352
}
343353

344354

@@ -374,18 +384,15 @@ SHA3_get_digest_size(PyObject *op, void *Py_UNUSED(closure))
374384
{
375385
// Preserving previous behavior: variable-length algorithms return 0
376386
SHA3object *self = _SHA3object_CAST(op);
377-
if (Hacl_Hash_SHA3_is_shake(self->hash_state))
378-
return PyLong_FromLong(0);
379-
else
380-
return PyLong_FromLong(Hacl_Hash_SHA3_hash_len(self->hash_state));
387+
return PyLong_FromLong(self->digest_size);
381388
}
382389

383390

384391
static PyObject *
385392
SHA3_get_capacity_bits(PyObject *op, void *Py_UNUSED(closure))
386393
{
387394
SHA3object *self = _SHA3object_CAST(op);
388-
uint32_t rate = Hacl_Hash_SHA3_block_len(self->hash_state) * 8;
395+
uint32_t rate = self->block_size * 8;
389396
assert(rate <= 1600);
390397
int capacity = 1600 - rate;
391398
return PyLong_FromLong(capacity);
@@ -396,8 +403,7 @@ static PyObject *
396403
SHA3_get_rate_bits(PyObject *op, void *Py_UNUSED(closure))
397404
{
398405
SHA3object *self = _SHA3object_CAST(op);
399-
uint32_t rate = Hacl_Hash_SHA3_block_len(self->hash_state) * 8;
400-
return PyLong_FromLong(rate);
406+
return PyLong_FromLong(self->block_size * 8);
401407
}
402408

403409
static PyObject *

0 commit comments

Comments
 (0)