Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 29 additions & 16 deletions src/python/pose_format/pose_header.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
25 changes: 25 additions & 0 deletions src/python/tests/pose_test.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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:
Expand Down
Loading