Skip to content

Commit cf0465e

Browse files
committed
fix key conflict issue
1 parent edfb54e commit cf0465e

4 files changed

Lines changed: 73 additions & 24 deletions

File tree

src/schematic/client.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -247,7 +247,7 @@ def _check_flag_via_api(
247247
return cached_value
248248

249249
resp = self.features.check_flag(flag_key, company=company, user=user)
250-
if resp is None or resp.data.value is None:
250+
if resp is None or resp.data is None or resp.data.value is None:
251251
return self._default_response(flag_key, options, REASON_FLAG_NOT_FOUND)
252252

253253
self._safe_cache_set(cache_key, resp.data)
@@ -627,7 +627,7 @@ def _ds_result_to_response(
627627
"""Convert a RulesengineCheckFlagResult (from DataStream) into the
628628
public CheckFlagResponseData shape."""
629629
entitlement = (
630-
FeatureEntitlement.model_validate(resp.entitlement.model_dump(mode="json"))
630+
FeatureEntitlement.model_validate(resp.entitlement.model_dump())
631631
if resp.entitlement is not None else None
632632
)
633633
return CheckFlagResponseData(
@@ -658,7 +658,7 @@ async def _check_flag_via_api(
658658
return cached_value
659659

660660
resp = await self.features.check_flag(flag_key, company=company, user=user)
661-
if resp is None or resp.data.value is None:
661+
if resp is None or resp.data is None or resp.data.value is None:
662662
return self._default_response(flag_key, options, REASON_FLAG_NOT_FOUND)
663663

664664
self._safe_cache_set(cache_key, resp.data)

src/schematic/datastream/__init__.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
from .datastream_client import DataStreamClient, DataStreamClientOptions
33
from .merge import deep_copy_company, deep_copy_user, partial_company, partial_user
44
from .rules_engine import RulesEngineClient
5-
from .types import DataStreamBaseReq, DataStreamError, DataStreamReq, DataStreamResp, EntityType, MessageType
5+
from .types import DataStreamBaseReq, DataStreamError, DataStreamReq, DataStreamResp, EntityType, KeyConflictError, MessageType
66
from .websocket_client import ClientOptions, DatastreamWSClient, convert_api_url_to_websocket_url
77

88
__all__ = [
@@ -25,6 +25,7 @@
2525
"DataStreamReq",
2626
"DataStreamResp",
2727
"EntityType",
28+
"KeyConflictError",
2829
"MessageType",
2930
# WebSocket client
3031
"ClientOptions",

src/schematic/datastream/datastream_client.py

Lines changed: 64 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
from ..cache import AsyncCacheProvider, AsyncLocalCache
1717
from .merge import partial_company, partial_user
1818
from .rules_engine import RulesEngineClient
19-
from .types import DataStreamBaseReq, DataStreamReq, DataStreamResp, EntityType, MessageType
19+
from .types import DataStreamBaseReq, DataStreamReq, DataStreamResp, EntityType, KeyConflictError, MessageType
2020
from .websocket_client import ClientOptions as WSClientOptions, DatastreamWSClient
2121

2222

@@ -363,7 +363,10 @@ async def get_all_flags(self) -> None:
363363
try:
364364
await asyncio.wait_for(asyncio.shield(self._pending_flags), timeout=RESOURCE_TIMEOUT_MS / 1000)
365365
except asyncio.TimeoutError:
366+
fut = self._pending_flags
366367
self._pending_flags = None
368+
if fut is not None and not fut.done():
369+
fut.set_result(False)
367370
raise TimeoutError("Timeout while waiting for flags data")
368371

369372
async def check_flag(
@@ -384,10 +387,20 @@ async def check_flag(
384387
cached_company: Optional[Any] = None
385388
cached_user: Optional[Any] = None
386389

387-
if needs_company:
388-
cached_company = await self._get_company_from_cache(company_keys) # type: ignore[arg-type]
389-
if needs_user:
390-
cached_user = await self._get_user_from_cache(user_keys) # type: ignore[arg-type]
390+
try:
391+
if needs_company:
392+
cached_company = await self._get_company_from_cache(company_keys) # type: ignore[arg-type]
393+
if needs_user:
394+
cached_user = await self._get_user_from_cache(user_keys) # type: ignore[arg-type]
395+
except KeyConflictError as exc:
396+
self._logger.warning("Key conflict for flag %s: %s", flag_key, exc)
397+
return RulesengineCheckFlagResult(
398+
value=flag.default_value,
399+
reason="key conflict",
400+
flag_key=flag.key,
401+
flag_id=flag.id,
402+
err=str(exc),
403+
)
391404

392405
# Replicator mode — evaluate with whatever is cached
393406
if self._replicator_mode:
@@ -670,35 +683,65 @@ async def _send_request(self, request: DataStreamReq) -> None:
670683
# ------------------------------------------------------------------
671684

672685
async def _get_company_from_cache(self, keys: Dict[str, str]) -> Optional[RulesengineCompany]:
686+
matched_id: Optional[str] = None
673687
for key, value in keys.items():
674688
ck = self._resource_key_to_cache_key(_PREFIX_COMPANY, key, value)
675689
try:
676690
company_id = await self._company_key_cache.get(ck)
677-
self._logger.debug("Company lookup key %s -> %s", ck, company_id)
678-
if company_id:
679-
rk = self._resource_id_cache_key(_PREFIX_COMPANY, company_id)
680-
raw = await self._company_cache.get(rk)
681-
self._logger.debug("Company ID key %s -> %s", rk, "hit" if raw is not None else "miss")
682-
if raw is not None:
683-
company = _validate(RulesengineCompany, raw)
684-
return company.model_copy(deep=True)
685691
except Exception as exc:
686692
self._logger.warning("Failed to retrieve company from cache: %s", exc)
693+
continue
694+
self._logger.debug("Company lookup key %s -> %s", ck, company_id)
695+
if not company_id:
696+
continue
697+
if matched_id is None:
698+
matched_id = company_id
699+
elif matched_id != company_id:
700+
raise KeyConflictError(
701+
f"Company keys match multiple entities: {matched_id} and {company_id}"
702+
)
703+
704+
if matched_id is None:
705+
return None
706+
707+
try:
708+
rk = self._resource_id_cache_key(_PREFIX_COMPANY, matched_id)
709+
raw = await self._company_cache.get(rk)
710+
self._logger.debug("Company ID key %s -> %s", rk, "hit" if raw is not None else "miss")
711+
if raw is not None:
712+
return _validate(RulesengineCompany, raw).model_copy(deep=True)
713+
except Exception as exc:
714+
self._logger.warning("Failed to retrieve company from cache: %s", exc)
687715
return None
688716

689717
async def _get_user_from_cache(self, keys: Dict[str, str]) -> Optional[RulesengineUser]:
718+
matched_id: Optional[str] = None
690719
for key, value in keys.items():
691720
ck = self._resource_key_to_cache_key(_PREFIX_USER, key, value)
692721
try:
693722
user_id = await self._user_key_cache.get(ck)
694-
if user_id:
695-
rk = self._resource_id_cache_key(_PREFIX_USER, user_id)
696-
raw = await self._user_cache.get(rk)
697-
if raw is not None:
698-
user = _validate(RulesengineUser, raw)
699-
return user.model_copy(deep=True)
700723
except Exception as exc:
701724
self._logger.warning("Failed to retrieve user from cache: %s", exc)
725+
continue
726+
if not user_id:
727+
continue
728+
if matched_id is None:
729+
matched_id = user_id
730+
elif matched_id != user_id:
731+
raise KeyConflictError(
732+
f"User keys match multiple entities: {matched_id} and {user_id}"
733+
)
734+
735+
if matched_id is None:
736+
return None
737+
738+
try:
739+
rk = self._resource_id_cache_key(_PREFIX_USER, matched_id)
740+
raw = await self._user_cache.get(rk)
741+
if raw is not None:
742+
return _validate(RulesengineUser, raw).model_copy(deep=True)
743+
except Exception as exc:
744+
self._logger.warning("Failed to retrieve user from cache: %s", exc)
702745
return None
703746

704747
async def _cache_company(self, company: RulesengineCompany) -> None:
@@ -956,7 +999,8 @@ def _make_default_result(
956999
def _start_replicator_health_check(self) -> None:
9571000
if not self._replicator_health_url:
9581001
return
959-
self._health_check_client = httpx.AsyncClient(timeout=REPLICATOR_HEALTH_TIMEOUT_S)
1002+
if self._health_check_client is None:
1003+
self._health_check_client = httpx.AsyncClient(timeout=REPLICATOR_HEALTH_TIMEOUT_S)
9601004
self._logger.info(
9611005
"Starting replicator health check: url=%s, interval=%dms",
9621006
self._replicator_health_url, self._replicator_health_check_ms,

src/schematic/datastream/types.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -70,3 +70,7 @@ class DataStreamError:
7070
error: str
7171
keys: Optional[Dict[str, str]] = None
7272
entity_type: Optional[EntityType] = None
73+
74+
75+
class KeyConflictError(Exception):
76+
"""Raised when lookup keys resolve to multiple distinct entities."""

0 commit comments

Comments
 (0)