Skip to content
Open
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
38 changes: 37 additions & 1 deletion src/schematic/datastream/datastream_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
import logging
import typing
from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional
from typing import Any, Callable, Dict, List, Optional, Tuple

import httpx
from ..cache import AsyncCacheProvider, AsyncLocalCache
Expand Down Expand Up @@ -263,6 +263,10 @@ def __init__(self, options: DataStreamClientOptions) -> None:
self._company_locks: Dict[str, asyncio.Lock] = {}
self._user_locks: Dict[str, asyncio.Lock] = {}

# Set after the first connection is ready, so later ones are known to
# be reconnects.
self._has_connected = False

# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
Expand Down Expand Up @@ -770,6 +774,9 @@ async def _handle_error_message(self, message: DataStreamResp) -> None:

async def _handle_connection_ready(self) -> None:
self._logger.info("DataStream connection is ready")
if self._has_connected:
await self._clear_entity_caches()
self._has_connected = True
try:
# Only send the flags request — don't await the response here.
# The response will be processed by the message loop, which hasn't
Expand All @@ -783,6 +790,35 @@ async def _handle_connection_ready(self) -> None:
self._pending_flags = None
raise

async def _clear_entity_caches(self) -> None:
"""Drop cached companies and users after a reconnect.

The server only pushes updates for entities a connection requested and
forgets those requests when the connection closes, so cached entities
would stop receiving updates. Clearing makes the next check refetch
over the new connection, which subscribes again. Re-requesting every
cached entity here instead would send a burst of lookups the moment a
deploy drops every connection.

Not done on disconnect: while disconnected, a cache miss falls back
to the REST API, so clearing then would send every check there.
"""
version = self._get_version_key()
caches: List[Tuple[AsyncCacheProvider[Any], str]] = [
(self._company_cache, _PREFIX_COMPANY),
(self._company_key_cache, _PREFIX_COMPANY),
(self._user_cache, _PREFIX_USER),
(self._user_key_cache, _PREFIX_USER),
]
for cache, prefix in caches:
try:
await cache.delete_missing([], scan_pattern=f"{prefix}:{version}:*")
except Exception as exc:
self._logger.warning(
"Failed to clear %s cache on reconnect; cached entries may go stale until they expire: %s",
prefix, exc,
)

# ------------------------------------------------------------------
# Request sending
# ------------------------------------------------------------------
Expand Down
159 changes: 159 additions & 0 deletions tests/datastream/test_datastream_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -917,6 +917,165 @@ async def test_get_company_times_out_when_no_response(
await client.get_company({"slug": "never-arrives"})


def _company_data(company_id: str, keys: Dict[str, str]) -> Dict[str, Any]:
return {
"id": company_id,
"keys": keys,
"account_id": "acc_1",
"environment_id": "env_1",
"billing_product_ids": [],
"credit_balances": {},
"metrics": [],
"plan_ids": [],
"plan_version_ids": [],
"rules": [],
"traits": [],
}


def _user_data(user_id: str, keys: Dict[str, str]) -> Dict[str, Any]:
return {
"id": user_id,
"keys": keys,
"account_id": "acc_1",
"environment_id": "env_1",
"rules": [],
"traits": [],
}


class TestDataStreamClientClearOnReconnect:
"""The server forgets which entities a connection asked for when it closes,
so a reconnect drops cached companies and users and lets the next check
fetch them again over the new connection."""

def _connected_client(self, logger: logging.Logger, **cache_options: Any) -> tuple[DataStreamClient, List[Any]]:
client = DataStreamClient(DataStreamClientOptions(
api_key="test-key",
base_url="https://api.schematichq.com",
logger=logger,
**cache_options,
))
sent: List[Any] = []
fake_ws = MagicMock()
fake_ws.is_connected = MagicMock(return_value=True)
fake_ws.send_message = AsyncMock(side_effect=lambda req: sent.append(req))
client._ws_client = fake_ws
return client, sent

async def _fetch_company(
self, client: DataStreamClient, sent: List[Any], keys: Dict[str, str], company_id: str,
) -> None:
async def respond() -> None:
for _ in range(50):
if any(r.data.entity_type == EntityType.COMPANY for r in sent):
break
await asyncio.sleep(0.01)
await client._handle_message(DataStreamResp(
data=_company_data(company_id, keys),
entity_type=EntityType.COMPANY.value,
message_type=MessageType.FULL.value,
))

responder = asyncio.create_task(respond())
try:
await client.get_company(keys)
finally:
await responder

async def _cache_user(self, client: DataStreamClient, user_id: str, keys: Dict[str, str]) -> None:
await client._handle_message(DataStreamResp(
data=_user_data(user_id, keys),
entity_type=EntityType.USER.value,
message_type=MessageType.FULL.value,
))

async def _reconnect(self, client: DataStreamClient, sent: List[Any]) -> None:
client._on_ws_disconnected()
sent.clear()
await client._handle_connection_ready()

async def test_reconnect_clears_cache_and_next_get_refetches(self, logger: logging.Logger) -> None:
client, sent = self._connected_client(logger)
await client._handle_connection_ready()
await self._fetch_company(client, sent, {"slug": "acme"}, "co_1")
await self._cache_user(client, "u_1", {"email": "a@b.co"})

await self._reconnect(client, sent)

assert [r.data.entity_type for r in sent] == [EntityType.FLAGS]
assert await client._get_company_from_cache({"slug": "acme"}) is None
assert await client._get_user_from_cache({"email": "a@b.co"}) is None
assert await client._company_cache.get(client._resource_id_cache_key("company", "co_1")) is None
assert await client._user_cache.get(client._resource_id_cache_key("user", "u_1")) is None

sent.clear()
await self._fetch_company(client, sent, {"slug": "acme"}, "co_1")
assert [(r.data.entity_type, r.data.keys) for r in sent] == [(EntityType.COMPANY, {"slug": "acme"})]

async def test_first_connect_keeps_prepopulated_cache(self, logger: logging.Logger) -> None:
client, sent = self._connected_client(logger)
await client._handle_message(DataStreamResp(
data=_company_data("co_1", {"slug": "acme"}),
entity_type=EntityType.COMPANY.value,
message_type=MessageType.FULL.value,
))
await self._cache_user(client, "u_1", {"email": "a@b.co"})

await client._handle_connection_ready()

assert await client._get_company_from_cache({"slug": "acme"}) is not None
assert await client._get_user_from_cache({"email": "a@b.co"}) is not None

async def test_disconnect_alone_keeps_cache(self, logger: logging.Logger) -> None:
client, sent = self._connected_client(logger)
await client._handle_connection_ready()
await self._fetch_company(client, sent, {"slug": "acme"}, "co_1")

client._on_ws_disconnected()

assert await client._get_company_from_cache({"slug": "acme"}) is not None

async def test_reconnect_with_provider_lacking_delete_missing_keeps_going(self, logger: logging.Logger) -> None:
cache: CacheProvider[Any] = MockCacheProvider()
cache.delete_missing = AsyncMock(side_effect=NotImplementedError) # type: ignore[method-assign]
client, sent = self._connected_client(logger, company_cache=cache, company_lookup_cache=cache)
await client._handle_connection_ready()

await self._reconnect(client, sent)

assert [r.data.entity_type for r in sent] == [EntityType.FLAGS]

async def test_reconnect_clears_redis_entities_but_not_flags(self, logger: logging.Logger) -> None:
import fakeredis.aioredis

from schematic.cache import RedisCache

redis = fakeredis.aioredis.FakeRedis()
cache: RedisCache[Any] = RedisCache(redis)
client, sent = self._connected_client(
logger,
company_cache=cache,
company_lookup_cache=cache,
user_cache=cache,
user_lookup_cache=cache,
flag_cache=cache,
)
await client._handle_connection_ready()
await self._fetch_company(client, sent, {"slug": "acme"}, "co_1")
await self._cache_user(client, "u_1", {"email": "a@b.co"})
await cache.set(client._flag_cache_key("my-flag"), {"key": "my-flag"})
await redis.set("schematic:credit-lease:co_1:cred_1", "lease")

await self._reconnect(client, sent)

remaining = sorted(k.decode() for k in await redis.keys("*"))
assert remaining == [
"schematic:credit-lease:co_1:cred_1",
f"schematic:{client._flag_cache_key('my-flag')}",
]


class TestDataStreamClientReplicatorHealthCheck:
"""Spec §Replicator Mode: health check polling against replicator_health_url."""

Expand Down
Loading