diff --git a/src/python/pose_format/pose_header.py b/src/python/pose_format/pose_header.py index 3640573..c334c3b 100644 --- a/src/python/pose_format/pose_header.py +++ b/src/python/pose_format/pose_header.py @@ -1,6 +1,7 @@ import hashlib import math import struct +import threading from typing import BinaryIO, List, Tuple, Optional, Union from .utils.reader import BufferReader, ConstStructs @@ -230,36 +231,45 @@ def __str__(self): class PoseHeaderCache: + # All access goes through _lock: an unsynchronized reader races with set_cache and + # can observe a half-updated cache (e.g. the new header with the old hash/offsets). + # check_cache therefore also returns end_offset, so callers position their reader + # from the same consistent snapshot instead of re-reading the class attribute. start_offset: int = None end_offset: int = None hash: str = None header: 'PoseHeader' = None + _lock = threading.Lock() @staticmethod def calc_hash(buffer: bytes): return hashlib.md5(buffer[PoseHeaderCache.start_offset:PoseHeaderCache.end_offset]).hexdigest() @staticmethod - def check_cache(buffer: bytes) -> 'PoseHeader': - if PoseHeaderCache.hash is None: - return None + def check_cache(buffer: bytes) -> Optional[Tuple['PoseHeader', int]]: + with PoseHeaderCache._lock: + if PoseHeaderCache.hash is None: + return None - if PoseHeaderCache.hash == PoseHeaderCache.calc_hash(buffer): - return PoseHeaderCache.header + if PoseHeaderCache.hash == PoseHeaderCache.calc_hash(buffer): + return PoseHeaderCache.header, PoseHeaderCache.end_offset + return None @staticmethod def clear_cache(): - PoseHeaderCache.start_offset = None - PoseHeaderCache.end_offset = None - PoseHeaderCache.hash = None - PoseHeaderCache.header = None + with PoseHeaderCache._lock: + PoseHeaderCache.start_offset = None + PoseHeaderCache.end_offset = None + PoseHeaderCache.hash = None + PoseHeaderCache.header = None @staticmethod def set_cache(header: 'PoseHeader', buffer: bytes, start_offset: int, end_offset: int): - PoseHeaderCache.start_offset = start_offset - PoseHeaderCache.end_offset = end_offset - PoseHeaderCache.header = header - PoseHeaderCache.hash = PoseHeaderCache.calc_hash(buffer) + with PoseHeaderCache._lock: + PoseHeaderCache.start_offset = start_offset + PoseHeaderCache.end_offset = end_offset + PoseHeaderCache.header = header + PoseHeaderCache.hash = PoseHeaderCache.calc_hash(buffer) class PoseHeader: @@ -317,9 +327,12 @@ def read(reader: BufferReader) -> 'PoseHeader': PoseHeader An instance of PoseHeader. """ - cached_header = PoseHeaderCache.check_cache(reader.buffer) - if cached_header is not None: - reader.read_offset = PoseHeaderCache.end_offset + cached = PoseHeaderCache.check_cache(reader.buffer) + if cached is not None: + # header and end_offset come from the same atomic cache snapshot -- reading + # PoseHeaderCache.end_offset here instead would race with concurrent set_cache + cached_header, end_offset = cached + reader.read_offset = end_offset return cached_header start_offset = reader.read_offset diff --git a/src/python/tests/pose_test.py b/src/python/tests/pose_test.py index 514e27d..a05841d 100644 --- a/src/python/tests/pose_test.py +++ b/src/python/tests/pose_test.py @@ -1,5 +1,6 @@ import random import string +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from typing import Optional, Tuple from unittest import TestCase @@ -340,6 +341,30 @@ def test_read_empty_pose_body_shape_matches_numpy_pose_body(self): self.assertEqual(empty_pose.body.confidence.shape, numpy_pose.body.confidence.shape) self.assertEqual(empty_pose.body.fps, numpy_pose.body.fps) + def test_concurrent_read_of_different_files_is_safe(self): + # Regression test: PoseHeaderCache used to be updated non-atomically, so + # concurrent Pose.read calls on files with different headers could crash + # ("buffer is too small for requested array") or, worse, silently return a + # pose parsed with another file's header. + data_dir = Path(__file__).parent / "data" + buffers = {} + expected = {} + for name in ["mediapipe.pose", "openpose.pose"]: + with open(data_dir / name, 'rb') as f: + buffers[name] = f.read() + pose = Pose.read(buffers[name]) + expected[name] = ([c.name for c in pose.header.components], pose.body.data.shape) + + def read_one(name): + pose = Pose.read(buffers[name]) + return name, ([c.name for c in pose.header.components], pose.body.data.shape) + + names = list(buffers.keys()) * 4 + with ThreadPoolExecutor(max_workers=8) as executor: + for _ in range(30): + for name, got in executor.map(read_one, names): + self.assertEqual(expected[name], got) + def test_pose_bbox(self): data_dir = Path(__file__).parent / "data" with open(data_dir / 'mediapipe.pose', 'rb') as f: