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
25 changes: 21 additions & 4 deletions proxy_workbench/apiv1.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
import threading
import time
from collections import OrderedDict, deque
from contextlib import contextmanager
from dataclasses import dataclass, field
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
Expand Down Expand Up @@ -826,12 +827,28 @@ def __init__(self, ttl_s, maximum, clock):
self._clock = clock
self._records = OrderedDict()
self._lock = threading.Lock()
# Serialize identical mutations from check through execution and cache
# insertion. A fixed set of locks keeps memory bounded for random keys.
self._singleflight = tuple(threading.RLock() for _ in range(64))
# Only identical mutations share a lock. Hash-striped locks also held
# unrelated requests, changing their concurrency-limit result.
self._flight_guard = threading.Lock()
self._flights = {}

@contextmanager
def serialized(self, bucket, key):
return self._singleflight[hash((bucket, key)) % len(self._singleflight)]
identity = (bucket, key)
with self._flight_guard:
entry = self._flights.get(identity)
if entry is None:
entry = {'lock': threading.RLock(), 'users': 0}
self._flights[identity] = entry
entry['users'] += 1
try:
with entry['lock']:
yield
finally:
with self._flight_guard:
entry['users'] -= 1
if not entry['users']:
self._flights.pop(identity, None)

def get(self, bucket, key, digest):
with self._lock:
Expand Down
22 changes: 22 additions & 0 deletions tests/test_apiv1_resilience.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
"""Regression tests for API validation, concurrency, streams and scoped reads."""
from concurrent.futures import ThreadPoolExecutor
import http.client
from itertools import count
import json
from pathlib import Path
import sqlite3
Expand Down Expand Up @@ -147,6 +148,27 @@ def test_concurrent_replay_executes_a_mutation_once(self):
self.assertEqual([answer.status_code for answer in answers], [200, 200])
self.assertEqual(self.service.calls, 1)

def test_unrelated_keys_do_not_share_a_hash_stripe(self):
store = apiv1.IdempotencyStore(60, 100, time.monotonic)
bucket = ('reader', 'POST', '/v1/collections')
first = 'first'
stripe = hash((bucket, first)) % 64
second = next(f'other-{index}' for index in count()
if hash((bucket, f'other-{index}')) % 64 == stripe)
acquired = threading.Event()

def enter_second():
with store.serialized(bucket, second):
acquired.set()

with store.serialized(bucket, first):
thread = threading.Thread(target=enter_second)
thread.start()
self.assertTrue(acquired.wait(2), 'different keys must run independently')
thread.join(2)
self.assertFalse(thread.is_alive())
self.assertFalse(store._flights)

def test_head_event_stream_releases_key_slot_and_sends_no_body(self):
control = apiv1.ApiV1(self.service, Keys(concurrency=1),
apiv1.ApiConfig(port=0))
Expand Down
Loading