Skip to content

Commit dd488a5

Browse files
committed
fix(flags): restore cold-worker remote cache fallback
1 parent d2cde46 commit dd488a5

5 files changed

Lines changed: 132 additions & 29 deletions

File tree

‎posthog/client.py‎

Lines changed: 4 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -898,7 +898,7 @@ def __init__(
898898
self.flag_fallback_cache_url = flag_fallback_cache_url
899899
self.flag_cache = self._initialize_flag_cache(flag_fallback_cache_url)
900900
self.flag_definition_version = 0
901-
# Without definitions, remote results have only process-local provenance.
901+
# Until definitions load, only explicitly remote Redis results are shared.
902902
self._flag_definition_fingerprint = f"remote-only:{uuid4().hex}"
903903
if self.flag_cache:
904904
self.flag_cache._advance_generation(
@@ -2910,7 +2910,6 @@ def _hash_flag_definitions(data: FlagDefinitionCacheData) -> str:
29102910
def _update_flag_state(
29112911
self,
29122912
data: FlagDefinitionCacheData,
2913-
old_flags_by_key: Optional[dict] = None,
29142913
*,
29152914
_fingerprint: Optional[str] = None,
29162915
) -> None:
@@ -2973,9 +2972,7 @@ def _load_feature_flags(self):
29732972
self.log.debug(
29742973
"[FEATURE FLAGS] Using cached flag definitions from external cache"
29752974
)
2976-
self._update_flag_state(
2977-
cached_data, old_flags_by_key=self.feature_flags_by_key or {}
2978-
)
2975+
self._update_flag_state(cached_data)
29792976
self._last_feature_flag_poll = datetime.now(tz=timezone.utc)
29802977
return
29812978
else:
@@ -3060,12 +3057,7 @@ def _fetch_feature_flags_from_api(self):
30603057
)
30613058
return
30623059

3063-
old_flags_by_key: dict[str, dict] = self.feature_flags_by_key or {}
3064-
self._update_flag_state(
3065-
response.data,
3066-
old_flags_by_key=old_flags_by_key,
3067-
_fingerprint=fingerprint,
3068-
)
3060+
self._update_flag_state(response.data, _fingerprint=fingerprint)
30693061

30703062
if self._flag_definition_cache_provider:
30713063
cache_data_to_store = {
@@ -3509,7 +3501,7 @@ def _get_feature_flag_result(
35093501
# The request-start generation is an invalidation boundary, not
35103502
# a claim about the server's definitions. Refresh rejects late writes.
35113503
if self.flag_cache and flag_result:
3512-
self.flag_cache.set_cached_flag(
3504+
self.flag_cache._set_cached_remote_flag(
35133505
distinct_id, key, flag_result, local_definition_version
35143506
)
35153507

‎posthog/test/test_flag_cache_publication.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -177,7 +177,7 @@ def write(value, version):
177177
def test_queued_old_write_cannot_replace_new_generation(client):
178178
cache = client.flag_cache
179179
old_version = client.flag_definition_version
180-
client._update_flag_state(definitions(2), client.feature_flags_by_key)
180+
client._update_flag_state(definitions(2))
181181
assert evaluate(client) is False
182182
cache.set_cached_flag("user", "person", "old result", old_version)
183183
assert cache.get_stale_cached_flag("user", "person").get_value() is False
@@ -231,7 +231,7 @@ def test_standalone_cache_can_reuse_invalidated_version(client):
231231

232232
def test_fork_replaces_held_cache_write_lock_and_preserves_fence(client):
233233
cache = client.flag_cache
234-
client._update_flag_state(definitions(2), client.feature_flags_by_key)
234+
client._update_flag_state(definitions(2))
235235
lock = cache._write_lock
236236
lock.acquire()
237237
try:

‎posthog/test/test_property_matching_version.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -231,7 +231,7 @@ def reload_during_evaluation(*args, **kwargs):
231231
nonlocal calls
232232
calls += 1
233233
if calls == 1:
234-
client._update_flag_state(definitions(2), client.feature_flags_by_key)
234+
client._update_flag_state(definitions(2))
235235
return original(*args, **kwargs)
236236

237237
with mock.patch(
@@ -481,7 +481,7 @@ def test_in_flight_result_is_not_cached_in_new_generation(client):
481481
original = client_module.match_feature_flag_properties
482482

483483
def reload_during_evaluation(*args, **kwargs):
484-
client._update_flag_state(definitions(2), client.feature_flags_by_key)
484+
client._update_flag_state(definitions(2))
485485
return original(*args, **kwargs)
486486

487487
with mock.patch(

‎posthog/test/test_redis_flag_cache_snapshot.py‎

Lines changed: 78 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -24,13 +24,13 @@ def workers():
2424
redis = FakeRedis()
2525
clients = []
2626

27-
def worker():
27+
def worker(secret_key="test-secret"):
2828
with mock.patch.object(
2929
Client, "_initialize_flag_cache", return_value=RedisFlagCache(redis)
3030
):
3131
client = Client(
3232
FAKE_TEST_API_KEY,
33-
secret_key="test-secret",
33+
secret_key=secret_key,
3434
send=False,
3535
enable_local_evaluation=False,
3636
)
@@ -253,7 +253,13 @@ def test_fork_retains_snapshot_binding(workers):
253253

254254

255255
def remote_success(client, during_request=None):
256-
details = FeatureFlag.from_json({"key": "person", "enabled": True})
256+
details = FeatureFlag.from_json(
257+
{
258+
"key": "person",
259+
"enabled": True,
260+
"metadata": {"payload": '{"source":"remote"}'},
261+
}
262+
)
257263

258264
def request(*args, **kwargs):
259265
if during_request:
@@ -274,8 +280,8 @@ def test_remote_only_success_remains_available_for_stale_fallback(workers):
274280
with mock.patch("posthog.client.get", side_effect=APIError(503, "offline")):
275281
assert remote_success(writer).get_value() is True
276282
assert fallback(writer).get_value() is True
277-
# Without verified definitions, a new Client cannot reuse this provenance.
278-
assert fallback(workers()) is None
283+
# Remote results remain usable by a worker starting during the outage.
284+
assert fallback(workers()).get_value() is True
279285

280286

281287
@pytest.mark.parametrize("transition", ["hydrate", "empty-hydrate", 401, 402])
@@ -299,6 +305,7 @@ def change():
299305
with mock.patch("posthog.client.get", side_effect=APIError(503, "offline")):
300306
assert remote_success(client, change).get_value() is True
301307
assert fallback(client) is None
308+
assert fallback(workers()) is None
302309
load(client, definitions(1))
303310
assert evaluate(client).get_value() is True
304311
assert fallback(client).get_value() is True
@@ -327,7 +334,7 @@ def test_reset_empty_fingerprint_cannot_cache_or_read_results(workers, status):
327334
assert fallback(client) is None
328335

329336

330-
def test_invalidated_remote_only_result_does_not_revive_after_restart(workers):
337+
def test_remote_result_requires_matching_snapshot_after_hydration(workers):
331338
writer = workers()
332339
with mock.patch("posthog.client.get", side_effect=APIError(503, "offline")):
333340
assert remote_success(writer).get_value() is True
@@ -336,10 +343,10 @@ def test_invalidated_remote_only_result_does_not_revive_after_restart(workers):
336343
assert fallback(writer) is None
337344
writer.shutdown()
338345
with mock.patch("posthog.client.get", side_effect=APIError(503, "offline")):
339-
assert fallback(workers()) is None
346+
assert fallback(workers()).get_value() is True
340347

341348

342-
def test_fork_renews_remote_only_provenance(workers):
349+
def test_fork_retains_remote_fallback(workers):
343350
client = workers()
344351
with mock.patch("posthog.client.get", side_effect=APIError(503, "offline")):
345352
assert remote_success(client).get_value() is True
@@ -348,6 +355,68 @@ def test_fork_renews_remote_only_provenance(workers):
348355
client, "_initialize_flag_cache", return_value=RedisFlagCache(redis)
349356
):
350357
client._reinit_after_fork()
351-
assert fallback(client) is None
358+
assert fallback(client).get_value() is True
352359
assert remote_success(client).get_value() is True
353360
assert fallback(client).get_value() is True
361+
362+
363+
@pytest.mark.parametrize("secret_key", [None, "test-secret"])
364+
@pytest.mark.parametrize("writer_has_definitions", [False, True])
365+
def test_cold_worker_uses_shared_remote_result_during_outage(
366+
workers, secret_key, writer_has_definitions
367+
):
368+
writer = workers(secret_key)
369+
if writer_has_definitions:
370+
load(writer, definitions(1))
371+
with mock.patch("posthog.client.get", side_effect=APIError(503, "offline")):
372+
assert remote_success(writer).get_value() is True
373+
reader = workers(secret_key)
374+
# Reproduce startup when neither endpoint is available.
375+
reader._load_feature_flags()
376+
assert reader.feature_flags is None
377+
result = fallback(reader)
378+
assert result is not None
379+
assert result.get_value() is True
380+
assert result.payload == {"source": "remote"}
381+
382+
383+
def test_cold_worker_rejects_shared_local_result_during_outage(workers):
384+
writer = workers()
385+
load(writer, definitions(1))
386+
assert evaluate(writer).get_value() is True
387+
with mock.patch("posthog.client.get", side_effect=APIError(503, "offline")):
388+
assert fallback(workers()) is None
389+
390+
391+
def test_loaded_worker_rejects_remote_result_from_different_snapshot(workers):
392+
writer = workers()
393+
load(writer, definitions(1))
394+
assert remote_success(writer).get_value() is True
395+
reader = workers()
396+
load(reader, definitions(2))
397+
assert fallback(reader) is None
398+
399+
400+
@pytest.mark.parametrize("status", [401, 402])
401+
def test_reset_worker_rejects_shared_remote_result(workers, status):
402+
writer = workers()
403+
load(writer, definitions(1))
404+
assert remote_success(writer).get_value() is True
405+
reader = workers()
406+
with mock.patch("posthog.client.get", side_effect=APIError(status, "reset")):
407+
reader._load_feature_flags()
408+
assert fallback(reader) is None
409+
410+
411+
def test_shared_remote_fallback_respects_stale_ttl(workers):
412+
writer = workers()
413+
with (
414+
mock.patch("posthog.client.get", side_effect=APIError(503, "offline")),
415+
mock.patch("posthog.utils.time.time", return_value=100) as clock,
416+
):
417+
assert remote_success(writer).get_value() is True
418+
reader = workers()
419+
clock.return_value = 100 + writer.flag_cache.stale_ttl - 1
420+
assert fallback(reader).get_value() is True
421+
clock.return_value += 1
422+
assert fallback(reader) is None

‎posthog/utils.py‎

Lines changed: 46 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -218,6 +218,7 @@ def __init__(self, flag_result, flag_definition_version, timestamp=None):
218218
self.flag_definition_version = flag_definition_version
219219
self.timestamp = timestamp or time.time()
220220
self._snapshot_fingerprint: Optional[str] = None
221+
self._is_remote = False
221222

222223
def is_valid(self, current_time, ttl, current_flag_version):
223224
time_valid = (current_time - self.timestamp) < ttl
@@ -237,6 +238,13 @@ def __init__(self, max_size=CACHE_MAX_SIZE, default_ttl=CACHE_TTL):
237238
self._minimum_version = None
238239
self._write_lock = threading.Lock()
239240

241+
def _set_cached_remote_flag(
242+
self, distinct_id, flag_key, flag_result, flag_definition_version
243+
):
244+
self.set_cached_flag(
245+
distinct_id, flag_key, flag_result, flag_definition_version
246+
)
247+
240248
def _advance_generation(self, version, fingerprint=None):
241249
# Client generations only move forward. Standalone cache invalidation
242250
# retains its existing exact-version deletion semantics.
@@ -385,6 +393,10 @@ def _advance_generation(self, version, fingerprint=None):
385393
def _is_entry_current(self, entry):
386394
snapshot = self._snapshot
387395
if snapshot is not None:
396+
# A worker starting during an outage has no local snapshot to verify.
397+
# Only explicitly remote results may cross that provenance boundary.
398+
if snapshot[1].startswith("remote-only:") and entry._is_remote:
399+
return True
388400
return bool(snapshot[1]) and entry._snapshot_fingerprint == snapshot[1]
389401
return self._is_version_current(entry.flag_definition_version)
390402

@@ -395,7 +407,12 @@ def _get_cache_key(self, distinct_id, flag_key):
395407
return f"{self.key_prefix}{distinct_id}:{flag_key}"
396408

397409
def _serialize_entry(
398-
self, flag_result, flag_definition_version, timestamp=None, fingerprint=None
410+
self,
411+
flag_result,
412+
flag_definition_version,
413+
timestamp=None,
414+
fingerprint=None,
415+
is_remote=False,
399416
):
400417
if timestamp is None:
401418
timestamp = time.time()
@@ -410,6 +427,8 @@ def _serialize_entry(
410427
}
411428
if fingerprint is not None:
412429
entry["snapshot_fingerprint"] = fingerprint
430+
if is_remote:
431+
entry["evaluation_source"] = "remote"
413432
if isinstance(flag_result, _FeatureFlagResult):
414433
# Additive metadata keeps the existing entry shape readable by older SDKs.
415434
entry["flag_result_type"] = _FEATURE_FLAG_RESULT_TYPE
@@ -433,6 +452,7 @@ def _deserialize_entry(self, data):
433452
timestamp=entry["timestamp"],
434453
)
435454
result._snapshot_fingerprint = entry.get("snapshot_fingerprint")
455+
result._is_remote = entry.get("evaluation_source") == "remote"
436456
return result
437457
except (json.JSONDecodeError, KeyError, TypeError, ValueError):
438458
# If deserialization fails, treat as cache miss
@@ -486,6 +506,25 @@ def get_stale_cached_flag(self, distinct_id, flag_key, max_stale_age=None):
486506

487507
def set_cached_flag(
488508
self, distinct_id, flag_key, flag_result, flag_definition_version
509+
):
510+
self._set_cached_flag(
511+
distinct_id, flag_key, flag_result, flag_definition_version
512+
)
513+
514+
def _set_cached_remote_flag(
515+
self, distinct_id, flag_key, flag_result, flag_definition_version
516+
):
517+
self._set_cached_flag(
518+
distinct_id, flag_key, flag_result, flag_definition_version, is_remote=True
519+
)
520+
521+
def _set_cached_flag(
522+
self,
523+
distinct_id,
524+
flag_key,
525+
flag_result,
526+
flag_definition_version,
527+
is_remote=False,
489528
):
490529
try:
491530
cache_key = self._get_cache_key(distinct_id, flag_key)
@@ -498,16 +537,19 @@ def set_cached_flag(
498537
return
499538
fingerprint = snapshot[1]
500539
serialized_entry = self._serialize_entry(
501-
flag_result, flag_definition_version, fingerprint=fingerprint
540+
flag_result,
541+
flag_definition_version,
542+
fingerprint=fingerprint,
543+
is_remote=is_remote,
502544
)
503545

504546
# Serialize writes so an old in-flight SETEX cannot overwrite a newer
505547
# result. Publication advances the fence without taking this lock.
506548
with self._write_lock:
507549
if not self._is_version_current(flag_definition_version):
508550
return
509-
# Late writes keep their original fingerprint. Other workers
510-
# reject them too, even if their local generation counters differ.
551+
# Late writes keep their original fingerprint, so workers with
552+
# different loaded definitions cannot accept them.
511553
self.redis.setex(cache_key, self.stale_ttl, serialized_entry)
512554
self.redis.set(self.version_key, flag_definition_version)
513555

0 commit comments

Comments
 (0)