Skip to content
Closed
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
4 changes: 2 additions & 2 deletions src/schematic/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -564,7 +564,7 @@ async def check_flag_with_entitlement(
await self._enqueue_flag_check_event(flag_key, resp, company, user)
return self._ds_result_to_response(flag_key, resp, options)
except Exception as e:
self.logger.debug(f"Datastream flag check failed ({e}), falling back to API")
self.logger.warning(f"Datastream flag check failed ({e}), falling back to API")

return await self._check_flag_via_api(flag_key, company, user, options)

Expand Down Expand Up @@ -594,7 +594,7 @@ async def check_flags(
results.append(self._ds_result_to_response(flag_key, resp, options))
return results
except Exception as e:
self.logger.debug(f"Datastream check_flags failed ({e}), falling back to bulk API")
self.logger.warning(f"Datastream check_flags failed ({e}), falling back to bulk API")

return await self._check_flags_via_api(flag_keys, company, user, options)

Expand Down
12 changes: 11 additions & 1 deletion src/schematic/datastream/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,16 @@
from .datastream_client import DataStreamClient, DataStreamClientOptions
from .merge import deep_copy_company, deep_copy_user, partial_company, partial_user
from .rules_engine import RulesEngineClient
from .types import DataStreamBaseReq, DataStreamError, DataStreamReq, DataStreamResp, EntityType, KeyConflictError, MessageType
from .types import (
DataStreamBaseReq,
DataStreamError,
DataStreamReq,
DataStreamResp,
EntityType,
KeyConflictError,
MessageType,
RulesEngineError,
)
from .websocket_client import ClientOptions, DatastreamWSClient, convert_api_url_to_websocket_url

__all__ = [
Expand All @@ -27,6 +36,7 @@
"EntityType",
"KeyConflictError",
"MessageType",
"RulesEngineError",
# WebSocket client
"ClientOptions",
"DatastreamWSClient",
Expand Down
39 changes: 13 additions & 26 deletions src/schematic/datastream/datastream_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
from ..cache import AsyncCacheProvider, AsyncLocalCache
from .merge import partial_company, partial_user
from .rules_engine import RulesEngineClient
from .types import DataStreamBaseReq, DataStreamReq, DataStreamResp, EntityType, KeyConflictError, MessageType
from .types import DataStreamBaseReq, DataStreamReq, DataStreamResp, EntityType, KeyConflictError, MessageType, RulesEngineError
from .websocket_client import ClientOptions as WSClientOptions, DatastreamWSClient


Expand Down Expand Up @@ -1003,34 +1003,21 @@ def _evaluate_flag(
company: Optional[RulesengineCompany],
user: Optional[RulesengineUser],
) -> RulesengineCheckFlagResult:
default_value = flag.default_value
"""Evaluate a flag with the local rules engine.

Raises ``RulesEngineError`` when the engine is unavailable or fails, so
the caller falls back to the REST API. Returning the flag's default here
would hand the caller a value indistinguishable from a real verdict.
"""
if not self._rules_engine.is_initialized():
self._logger.warning("Rules engine not initialized; flag %s cannot be evaluated locally", flag.key)
raise RulesEngineError(f"Rules engine not initialized (flag {flag.key})")

try:
if self._rules_engine.is_initialized():
return self._rules_engine.check_flag(flag, company, user)
else:
self._logger.warning("Rules engine not initialized, using default flag value")
return self._make_default_result(flag, company, user, default_value, "RULES_ENGINE_UNAVAILABLE")
return self._rules_engine.check_flag(flag, company, user)
except Exception as exc:
self._logger.error("Rules engine evaluation failed: %s", exc)
return self._make_default_result(flag, company, user, default_value, "RULES_ENGINE_ERROR")

@staticmethod
def _make_default_result(
flag: RulesengineFlag,
company: Optional[RulesengineCompany],
user: Optional[RulesengineUser],
value: bool,
reason: str,
) -> RulesengineCheckFlagResult:
return RulesengineCheckFlagResult(
value=value,
reason=reason,
flag_key=flag.key,
flag_id=flag.id,
company_id=company.id if company else None,
user_id=user.id if user else None,
)
self._logger.warning("Rules engine evaluation failed for flag %s: %s", flag.key, exc)
raise RulesEngineError(f"Rules engine evaluation failed for flag {flag.key}: {exc}") from exc

# ------------------------------------------------------------------
# Replicator health checking
Expand Down
25 changes: 22 additions & 3 deletions src/schematic/datastream/rules_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,25 @@ def _deep_camel_to_snake(obj: Any) -> Any:
return [_deep_camel_to_snake(item) for item in obj]
return obj

def _strip_none(obj: Any) -> Any:
"""Recursively drop dict entries whose value is None.

The generated Pydantic models keep explicitly-set None values even with
``exclude_none=True``, and the partial-update merge in ``merge.py`` sets
every unset optional field to None. The rules engine treats an absent key
and an explicit null differently: ``#[serde(default)]`` covers the former
only, so a null for a collection field rejects the whole envelope and the
check fails with an error code (schematichq 1.3.4 / WASM v0.7.0). Sending
only the keys that carry a value makes the envelope shape independent of
how the models were built.
"""
if isinstance(obj, dict):
return {k: _strip_none(v) for k, v in obj.items() if v is not None}
if isinstance(obj, list):
return [_strip_none(item) for item in obj]
return obj


# Path to the WASM binary shipped alongside this module
_WASM_PATH = Path(__file__).parent / "wasm" / "rulesengine.wasm"

Expand Down Expand Up @@ -140,9 +159,9 @@ def check_flag(
self._ensure_initialized()

envelope = {
"flag": flag.model_dump(exclude_none=True, mode="json"),
"company": company.model_dump(exclude_none=True, mode="json") if company else None,
"user": user.model_dump(exclude_none=True, mode="json") if user else None,
"flag": _strip_none(flag.model_dump(exclude_none=True, mode="json")),
"company": _strip_none(company.model_dump(exclude_none=True, mode="json")) if company else None,
"user": _strip_none(user.model_dump(exclude_none=True, mode="json")) if user else None,
}

result_json = self._call_wasm(json.dumps(envelope))
Expand Down
10 changes: 10 additions & 0 deletions src/schematic/datastream/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,3 +74,13 @@ class DataStreamError:

class KeyConflictError(Exception):
"""Raised when lookup keys resolve to multiple distinct entities."""


class RulesEngineError(Exception):
"""Raised when the local rules engine cannot evaluate a flag.

The datastream client raises this instead of returning the flag's default
value so that callers (the Schematic client) fall back to the REST API. A
failed evaluation must never be mistaken for a real verdict.
"""

50 changes: 43 additions & 7 deletions tests/datastream/test_datastream_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@

from schematic.cache import AsyncCacheProvider as CacheProvider, AsyncLocalCache as LocalCache
from schematic.datastream.datastream_client import DataStreamClient, DataStreamClientOptions
from schematic.datastream.types import DataStreamResp, EntityType, MessageType
from schematic.datastream.types import DataStreamResp, EntityType, MessageType, RulesEngineError
from schematic.types import CheckFlagRequestBody, RulesengineCheckFlagResult


Expand Down Expand Up @@ -318,7 +318,7 @@ def test_resource_key_to_cache_key_lowercases(self, logger: logging.Logger) -> N


class TestDataStreamClientFlagEvaluation:
async def test_evaluate_flag_returns_default_when_engine_unavailable(self, logger: logging.Logger) -> None:
async def test_evaluate_flag_raises_when_engine_unavailable(self, logger: logging.Logger) -> None:
from schematic.types import RulesengineFlag

cache = MockCacheProvider()
Expand All @@ -336,11 +336,41 @@ async def test_evaluate_flag_returns_default_when_engine_unavailable(self, logge
id="f1", key="test", account_id="a", environment_id="e",
default_value=True, rules=[],
)
result = client._evaluate_flag(flag, None, None)
assert isinstance(result, RulesengineCheckFlagResult)
assert result.value is True
assert result.reason == "RULES_ENGINE_UNAVAILABLE"
assert result.flag_key == "test"
# An uninitialized engine must not produce a value that looks like a
# verdict; raising lets the Schematic client fall back to the API.
with pytest.raises(RulesEngineError, match="not initialized"):
client._evaluate_flag(flag, None, None)

async def test_evaluate_flag_raises_when_engine_fails(self, logger: logging.Logger) -> None:
"""A rules engine failure must surface as an exception, not as the flag default.

Regression for schematichq 1.3.4: the WASM rejected the envelope and the
client returned the flag default with reason RULES_ENGINE_ERROR, which
callers could not distinguish from a genuine "not entitled" answer.
"""
from schematic.types import RulesengineFlag

cache = MockCacheProvider()
client = DataStreamClient(DataStreamClientOptions(
api_key="test-key",
logger=logger,
replicator_mode=True,
company_cache=cache,
company_lookup_cache=cache,
user_cache=cache,
user_lookup_cache=cache,
flag_cache=cache,
))
flag = RulesengineFlag(
id="f1", key="test", account_id="a", environment_id="e",
default_value=False, rules=[],
)
client._rules_engine = MagicMock()
client._rules_engine.is_initialized.return_value = True
client._rules_engine.check_flag.side_effect = RuntimeError("WASM checkFlagCombined returned error code")

with pytest.raises(RulesEngineError, match="returned error code"):
client._evaluate_flag(flag, None, None)

async def test_check_flag_raises_when_flag_not_found(self, logger: logging.Logger) -> None:
cache = MockCacheProvider()
Expand Down Expand Up @@ -372,6 +402,9 @@ async def test_flag_evaluation_with_cached_company(self, logger: logging.Logger)
user_lookup_cache=cache,
flag_cache=cache,
))
# The engine used to hand back the flag default when uninitialized;
# it now raises, so evaluate with the real WASM.
await client._rules_engine.initialize()

# Cache a company via full message
await client._handle_message(DataStreamResp(
Expand Down Expand Up @@ -422,6 +455,9 @@ async def test_flag_evaluation_with_cached_user(self, logger: logging.Logger) ->
user_lookup_cache=cache,
flag_cache=cache,
))
# The engine used to hand back the flag default when uninitialized;
# it now raises, so evaluate with the real WASM.
await client._rules_engine.initialize()

# Cache a user
await client._handle_message(DataStreamResp(
Expand Down
83 changes: 83 additions & 0 deletions tests/datastream/test_rules_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -244,3 +244,86 @@ async def test_missing_wasm_raises(self) -> None:
engine = RulesEngineClient(wasm_path="/nonexistent/rulesengine.wasm")
with pytest.raises(FileNotFoundError):
await engine.initialize()


class TestRulesEngineEnvelopeNulls:
"""Regression for schematichq 1.3.4 / rules engine WASM v0.7.0.

The generated models keep explicitly-set ``None`` values when dumped with
``exclude_none=True``, and ``partial_company`` sets every unset optional
entitlement field to ``None`` when it merges a partial update. The WASM
treats an absent key and an explicit ``null`` differently, and rejected
``"warning_tiers": null`` with an error code, so every check against a
company failed from its first partial update onward. The envelope must
never carry nulls, whatever shape the models are in.
"""

@pytest.fixture
async def engine(self) -> RulesEngineClient:
e = RulesEngineClient()
await e.initialize()
return e

def _merged_company(self) -> RulesengineCompany:
from schematic.datastream.datastream_client import _validate
from schematic.datastream.merge import partial_company

raw = {
"id": "co_1",
"account_id": "acc_1",
"environment_id": "env_1",
"keys": {"id": "c1"},
"traits": [],
"metrics": [],
"rules": [],
"plan_ids": ["plan_1"],
"plan_version_ids": [],
"billing_product_ids": [],
"credit_balances": {},
"entitlements": [
{"feature_id": "feat_1", "feature_key": "test-flag", "value_type": "boolean"},
],
}
full = _validate(RulesengineCompany, raw)
return partial_company(full, {"credit_balances": {"crd_1": 5.0}})

def test_merged_company_dump_carries_explicit_nulls(self) -> None:
# Documents the model behaviour the envelope has to defend against. If
# this ever starts failing, the stripping below is no longer load-bearing.
dumped = self._merged_company().model_dump(exclude_none=True, mode="json")
assert "warning_tiers" in dumped["entitlements"][0]
assert dumped["entitlements"][0]["warning_tiers"] is None

async def test_envelope_contains_no_nulls(self, engine: RulesEngineClient) -> None:
import json

captured: list[str] = []
original = engine._call_wasm

def spy(input_json: str) -> str:
captured.append(input_json)
return original(input_json)

engine._call_wasm = spy # type: ignore[method-assign]
engine.check_flag(_make_flag(default_value=True), self._merged_company())

assert len(captured) == 1
envelope = json.loads(captured[0])
assert envelope["user"] is None # top-level absence is still expressed as null

def has_null(obj: object) -> bool:
if isinstance(obj, dict):
return any(v is None or has_null(v) for v in obj.values())
if isinstance(obj, list):
return any(item is None or has_null(item) for item in obj)
return False

assert not has_null(envelope["flag"])
assert not has_null(envelope["company"])
assert "warning_tiers" not in envelope["company"]["entitlements"][0]

async def test_check_flag_after_partial_merge_evaluates(self, engine: RulesEngineClient) -> None:
result = engine.check_flag(_make_flag(default_value=True), self._merged_company())
assert isinstance(result, RulesengineCheckFlagResult)
assert result.value is True
assert result.err is None
Loading