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
37 changes: 27 additions & 10 deletions ldclient/impl/async_big_segments.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import asyncio
from typing import Awaitable, Callable, Optional, Tuple

from expiringdict import ExpiringDict
Expand Down Expand Up @@ -61,6 +62,7 @@ def __init__(self, config: AsyncBigSegmentsConfig):
self.__stale_after_millis = config.stale_after * 1000
self.__status_provider = AsyncBigSegmentStoreStatusProviderImpl(self.get_status)
self.__last_status = None # type: Optional[BigSegmentStoreStatus]
self.__poll_lock = asyncio.Lock()
self.__poll_task = None # type: Optional[AsyncRepeatingTask]

if self.__store:
Expand Down Expand Up @@ -97,26 +99,37 @@ async def get_user_membership(self, user_key: str) -> Tuple[Optional[dict], str]
except Exception as e:
log.exception("Big Segment store membership query returned error: %s" % e)
return None, BigSegmentsStatus.STORE_ERROR
# First-call fallback: if the polling task hasn't run yet, poll inline now
status = self.__last_status
if status is None:
status = await self.poll_store_and_update_status()
status = await self.get_status()
if not status.available:
return membership, BigSegmentsStatus.STORE_ERROR
return membership, BigSegmentsStatus.STALE if status.stale else BigSegmentsStatus.HEALTHY

async def get_status(self) -> BigSegmentStoreStatus:
"""Return the most recently polled status.

When no status has been cached yet, poll the store inline and wait for the
result, so the status is accurate even if called immediately after start().
When no status has been cached yet, poll the store and wait for the result,
so the status is accurate even if called immediately after start().
"""
status = self.__last_status
if status is None:
status = await self.poll_store_and_update_status()
return status
if status is not None:
return status
# Check again under the lock: another caller may have polled while we waited.
async with self.__poll_lock:
status = self.__last_status
if status is not None:
return status
new_status = await self.__query_store_status()
self.__notify_status(new_status)
return new_status

async def poll_store_and_update_status(self) -> BigSegmentStoreStatus:
async with self.__poll_lock:
new_status = await self.__query_store_status()
self.__notify_status(new_status)
return new_status

async def __query_store_status(self) -> BigSegmentStoreStatus:
"""Queries the store and caches the result. Callers must hold the poll lock."""
new_status = BigSegmentStoreStatus(False, False) # default to "unavailable" if we don't get a new status below
if self.__store:
try:
Expand All @@ -125,8 +138,12 @@ async def poll_store_and_update_status(self) -> BigSegmentStoreStatus:
except Exception as e:
log.exception("Big Segment store status query returned error: %s" % e)
self.__last_status = new_status
self.__status_provider._update_status(new_status)
return new_status

def __notify_status(self, new_status: BigSegmentStoreStatus):
"""Tells the status provider about a new status, outside the poll lock so
that a listener calling back into the manager cannot deadlock."""
self.__status_provider._update_status(new_status)

def is_stale(self, timestamp) -> bool:
return is_stale(timestamp, self.__stale_after_millis)
30 changes: 25 additions & 5 deletions ldclient/impl/big_segments.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
from threading import Lock
from typing import Callable, Optional, Tuple

from expiringdict import ExpiringDict
Expand Down Expand Up @@ -63,6 +64,7 @@ def __init__(self, config: BigSegmentsConfig):
self.__stale_after_millis = config.stale_after * 1000
self.__status_provider = BigSegmentStoreStatusProviderImpl(self.get_status)
self.__last_status = None # type: Optional[BigSegmentStoreStatus]
self.__poll_lock = Lock()
self.__poll_task = None # type: Optional[RepeatingTask]

if self.__store:
Expand Down Expand Up @@ -94,18 +96,32 @@ def get_user_membership(self, user_key: str) -> Tuple[Optional[dict], str]:
except Exception as e:
log.exception("Big Segment store membership query returned error: %s" % e)
return None, BigSegmentsStatus.STORE_ERROR
status = self.__last_status
if not status:
status = self.poll_store_and_update_status()
status = self.get_status()
if not status.available:
return membership, BigSegmentsStatus.STORE_ERROR
return membership, BigSegmentsStatus.STALE if status.stale else BigSegmentsStatus.HEALTHY

def get_status(self) -> BigSegmentStoreStatus:
status = self.__last_status
return status if status else self.poll_store_and_update_status()
if status is not None:
return status
# Check again under the lock: another caller may have polled while we waited.
with self.__poll_lock:
status = self.__last_status
if status is not None:
return status
new_status = self.__query_store_status()
self.__notify_status(new_status)
return new_status

def poll_store_and_update_status(self) -> BigSegmentStoreStatus:
with self.__poll_lock:
new_status = self.__query_store_status()
self.__notify_status(new_status)
return new_status

def __query_store_status(self) -> BigSegmentStoreStatus:
"""Queries the store and caches the result. Callers must hold the poll lock."""
new_status = BigSegmentStoreStatus(False, False) # default to "unavailable" if we don't get a new status below
if self.__store:
try:
Expand All @@ -114,8 +130,12 @@ def poll_store_and_update_status(self) -> BigSegmentStoreStatus:
except Exception as e:
log.exception("Big Segment store status query returned error: %s" % e)
self.__last_status = new_status
self.__status_provider._update_status(new_status)
return new_status

def __notify_status(self, new_status: BigSegmentStoreStatus):
"""Tells the status provider about a new status, outside the poll lock so
that a listener calling back into the manager cannot deadlock."""
self.__status_provider._update_status(new_status)

def is_stale(self, timestamp) -> bool:
return is_stale(timestamp, self.__stale_after_millis)
41 changes: 41 additions & 0 deletions ldclient/testing/impl/test_async_big_segments.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,10 +26,17 @@ class MockAsyncBigSegmentStore(AsyncBigSegmentStore):

def __init__(self):
self._membership_queries = []
self._metadata_queries = []
self._memberships = {}
self._metadata_fn = lambda: BigSegmentStoreMetadata(int(time.time() * 1000))
self._metadata_delay = 0.0
self._stopped = False

def setup_metadata_delay(self, delay: float):
"""Makes get_metadata await for the given number of seconds, so a test can
keep a metadata query in flight."""
self._metadata_delay = delay

def setup_membership(self, user_hash: str, membership):
self._memberships[user_hash] = membership

Expand All @@ -48,6 +55,9 @@ def _raise():
self._metadata_fn = _raise

async def get_metadata(self) -> BigSegmentStoreMetadata:
self._metadata_queries.append(True)
if self._metadata_delay:
await asyncio.sleep(self._metadata_delay)
return self._metadata_fn()

async def get_membership(self, context_hash: str):
Expand All @@ -61,6 +71,10 @@ async def stop(self):
def membership_queries(self):
return list(self._membership_queries)

@property
def metadata_queries(self):
return list(self._metadata_queries)


async def make_started_manager(store, **kwargs):
config = AsyncBigSegmentsConfig(store=store, **kwargs)
Expand Down Expand Up @@ -348,3 +362,30 @@ async def test_get_status_with_no_store_configured():
assert status.available is False
finally:
await manager.stop()


@pytest.mark.asyncio
async def test_status_query_reuses_a_poll_that_is_already_in_flight():
"""
The polling task queries the store as soon as it starts. A status request that
arrives while that query is still in flight reuses it rather than sending a second
one, which is what the SDK does at startup.
"""
store = MockAsyncBigSegmentStore()
store.setup_metadata_always_up_to_date()
store.setup_metadata_delay(0.25)

manager = await make_started_manager(store, status_poll_interval=10)
try:
# Let the polling task start its query, so this request arrives while that
# query is in flight and must wait for it instead of starting its own.
await asyncio.sleep(0)
assert store.metadata_queries, "polling task never queried the store"

status = await manager.get_status()
assert status.available is True
await asyncio.sleep(0.1) # let a second query show up, if the fix is not working
finally:
await manager.stop()

assert len(store.metadata_queries) == 1
30 changes: 30 additions & 0 deletions ldclient/testing/impl/test_big_segments.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import time
from queue import Queue
from threading import Event

from ldclient.config import BigSegmentsConfig
from ldclient.evaluation import BigSegmentsStatus
Expand Down Expand Up @@ -185,3 +186,32 @@ def test_status_polling_detects_stale_status():
assert status3.stale is False
finally:
manager.stop()


def test_status_query_reuses_a_poll_that_is_already_in_flight():
# The polling task queries the store as soon as it starts. A status request
# that arrives while that query is still in flight reuses it rather than
# sending a second one, which is what the SDK does at startup.
metadata_queries = []
poll_started = Event()

def slow_metadata():
metadata_queries.append(True)
poll_started.set()
time.sleep(0.25)
return BigSegmentStoreMetadata(time.time() * 1000)

store = MockBigSegmentStore()
store.setup_metadata(slow_metadata)

manager = BigSegmentStoreManager(BigSegmentsConfig(store=store, status_poll_interval=10))
try:
assert poll_started.wait(1.0), "polling task never queried the store"
# The task's query is in flight now, so this request must wait for it
# instead of starting its own.
assert manager.status_provider.status.available is True
time.sleep(0.1) # let a second query show up, if the fix is not working
finally:
manager.stop()

assert len(metadata_queries) == 1
Loading