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
11 changes: 9 additions & 2 deletions src/schematic/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -1907,14 +1907,21 @@ async def track(
options=options,
)

# Update company metrics in DataStream if available and connected
# Bump the cached company metric so the next check sees this usage
await self._update_company_metrics(company, event, quantity)

async def _update_company_metrics(
self, company: Optional[Dict[str, str]], event: str, quantity: Optional[int],
) -> None:
# Not gated on is_connected(). Flag checks keep evaluating from the
# cache while a replicator reports not ready or the WebSocket is down,
# so the cached metric has to keep counting the usage tracked in the
# meantime or those checks would enforce limits against a frozen
# figure. The bump cannot double count: the server's next push for the
# company replaces the metric outright, and a company missing from the
# cache is left alone.
ds = self._get_datastream()
if company and ds is not None and ds.is_connected():
if company and ds is not None:
try:
await ds.update_company_metrics(
company,
Expand Down
17 changes: 17 additions & 0 deletions tests/custom/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -878,6 +878,23 @@ async def test_track(self):
)
mock_push.assert_called_once()

async def test_track_updates_company_metrics_when_datastream_not_connected(self):
"""The cached metric keeps counting while DataStream reports not
connected, since flag checks keep evaluating from that cache."""
mock_ds = MagicMock()
mock_ds.is_connected = MagicMock(return_value=False)
mock_ds.update_company_metrics = AsyncMock()
self.async_schematic._datastream_client = mock_ds

with patch.object(self.async_schematic.event_buffer, "push", new=AsyncMock()) as mock_push:
await self.async_schematic.track(
event="api-calls",
company={"id": "company_id"},
quantity=3,
)
mock_push.assert_awaited_once()
mock_ds.update_company_metrics.assert_awaited_once_with({"id": "company_id"}, "api-calls", 3)

async def test_track_with_options(self):
"""All TrackOptions fields must plumb through async track() to the
CreateEventRequestBody."""
Expand Down
196 changes: 196 additions & 0 deletions tests/custom/test_replicator_track.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,196 @@
"""Track's local usage update in replicator mode while the replicator reports
not ready.

When a Schematic account is closed or Schematic is unreachable, the replicator
stays up and keeps its cache but reports ``ready: false``. Flag checks keep
evaluating from that cache, so track has to keep bumping the cached company
metric, or a numeric limit would never trip while the replicator is down.
"""

from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch

from httpx import AsyncClient
from lease_support import make_fake_redis

from schematic.cache import RedisCache
from schematic.client import AsyncSchematic, AsyncSchematicConfig, DataStreamConfig
from schematic.types import (
RulesengineCompany,
RulesengineCompanyMetric,
RulesengineCondition,
RulesengineFlag,
RulesengineRule,
)

COMPANY_ID = "co_metered"
FLAG_KEY = "api-access"
EVENT = "api-calls"
LIMIT = 100


def _metered_company(usage: int) -> RulesengineCompany:
# A company override grants the flag while usage stays under LIMIT.
company_condition = RulesengineCondition(
id="cond_company",
account_id="acc_1",
environment_id="env_1",
condition_type="company",
operator="eq",
resource_ids=[COMPANY_ID],
trait_value="",
)
metric_condition = RulesengineCondition(
id="cond_metric",
account_id="acc_1",
environment_id="env_1",
condition_type="metric",
operator="lt",
resource_ids=[],
event_subtype=EVENT,
metric_value=LIMIT,
metric_period="all_time",
trait_value=str(LIMIT),
)
rule = RulesengineRule(
id="rule_override",
flag_id="flag_1",
account_id="acc_1",
environment_id="env_1",
name="Company Override",
rule_type="company_override",
value=True,
priority=0,
conditions=[company_condition, metric_condition],
condition_groups=[],
)
metric = RulesengineCompanyMetric(
account_id="acc_1",
environment_id="env_1",
company_id=COMPANY_ID,
event_subtype=EVENT,
period="all_time",
month_reset="first_of_month",
value=usage,
created_at="2026-01-01T00:00:00Z",
)
return RulesengineCompany(
id=COMPANY_ID,
account_id="acc_1",
environment_id="env_1",
keys={"id": COMPANY_ID},
traits=[],
metrics=[metric],
rules=[rule],
entitlements=[],
billing_product_ids=[],
credit_balances={},
plan_ids=[],
plan_version_ids=[],
)


def _flag() -> RulesengineFlag:
return RulesengineFlag(
id="flag_1",
key=FLAG_KEY,
account_id="acc_1",
environment_id="env_1",
default_value=False,
rules=[],
)


def _health_client(body: dict) -> MagicMock:
resp = MagicMock()
resp.raise_for_status = MagicMock()
resp.json = MagicMock(return_value=body)
client = MagicMock()
client.get = AsyncMock(return_value=resp)
return client


async def _replicator_client_not_ready(usage: int) -> AsyncSchematic:
"""An AsyncSchematic in replicator mode whose replicator has just reported
``ready: false`` with a metered company and its flag still in Redis."""
redis: RedisCache[Any] = RedisCache(make_fake_redis())
client = AsyncSchematic("test_key", AsyncSchematicConfig(
logger=MagicMock(),
httpx_client=MagicMock(spec=AsyncClient),
event_buffer_period=1,
flag_defaults={FLAG_KEY: False},
use_datastream=True,
datastream=DataStreamConfig(
replicator_mode=True,
replicator_health_url="http://replicator.test/ready",
company_cache=redis,
company_lookup_cache=redis,
user_cache=redis,
user_lookup_cache=redis,
flag_cache=redis,
),
))
ds = client._datastream_client
assert ds is not None
await ds._rules_engine.initialize()

# The replicator was ready, then reports not ready. The poll is what the
# background health loop runs, driven once here.
ds._replicator_ready = True
ds._health_check_client = _health_client({"ready": False, "cache_version": "v1"})
await ds._check_replicator_health()
assert not ds.is_connected()

await ds._cache_company(_metered_company(usage))
await ds._flag_cache.set(ds._flag_cache_key(FLAG_KEY), _flag())

client.features.check_flag = AsyncMock(side_effect=RuntimeError("api unreachable"))
client.flag_check_cache_providers = []
return client


async def test_track_counts_usage_locally_when_replicator_not_ready() -> None:
client = await _replicator_client_not_ready(usage=95)
ds = client._datastream_client
assert ds is not None
try:
with patch.object(client.event_buffer, "push", new=AsyncMock()) as push:
assert await client.check_flag(FLAG_KEY, company={"id": COMPANY_ID}) is True

await client.track(EVENT, company={"id": COMPANY_ID}, quantity=10)

# The event still goes out for the server to count.
track_events = [c.args[0] for c in push.await_args_list if c.args[0].event_type == "track"]
assert len(track_events) == 1

cached = await ds.get_cached_company({"id": COMPANY_ID})
assert cached is not None
assert cached.metrics is not None
assert cached.metrics[0].value == 105

# The cached figure now crosses the limit, so the check denies without
# waiting for the replicator to come back.
assert await client.check_flag(FLAG_KEY, company={"id": COMPANY_ID}) is False
client.features.check_flag.assert_not_called()
finally:
await client.shutdown()


async def test_track_for_an_uncached_company_is_a_no_op_when_replicator_not_ready() -> None:
client = await _replicator_client_not_ready(usage=0)
ds = client._datastream_client
assert ds is not None
try:
with patch.object(client.event_buffer, "push", new=AsyncMock()) as push:
await client.track(EVENT, company={"id": "co_unknown"}, quantity=10)
push.assert_awaited_once()

# Nothing is invented for a company the replicator never cached, and
# the cached one is untouched.
assert await ds.get_cached_company({"id": "co_unknown"}) is None
cached = await ds.get_cached_company({"id": COMPANY_ID})
assert cached is not None
assert cached.metrics is not None
assert cached.metrics[0].value == 0
finally:
await client.shutdown()
Loading