diff --git a/proxy_workbench/apiv1.py b/proxy_workbench/apiv1.py index a8c821e..044733e 100644 --- a/proxy_workbench/apiv1.py +++ b/proxy_workbench/apiv1.py @@ -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 @@ -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: diff --git a/tests/test_apiv1_resilience.py b/tests/test_apiv1_resilience.py index 1f89292..9aa7325 100644 --- a/tests/test_apiv1_resilience.py +++ b/tests/test_apiv1_resilience.py @@ -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 @@ -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))