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
7 changes: 5 additions & 2 deletions openfeature/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -890,7 +890,7 @@ async def _create_provider_evaluation_async(
if err := _typecheck_flag_value(value=resolution.value, flag_type=flag_type):
return FlagEvaluationDetails(
flag_key=flag_key,
value=resolution.value,
value=default_value,
reason=Reason.ERROR,
error_code=err.error_code,
error_message=err.error_message,
Expand Down Expand Up @@ -1001,6 +1001,9 @@ def _typecheck_flag_value(
py_type = type_map.get(flag_type)
if not py_type:
return GeneralError(error_message="Unknown flag type")
if not isinstance(value, py_type):
# bool subclasses int, but booleans are not integer flag values.
if not isinstance(value, py_type) or (
flag_type == FlagType.INTEGER and isinstance(value, bool)
):
return TypeMismatchError(f"Expected type {py_type} but got {type(value)}")
return None
117 changes: 117 additions & 0 deletions tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -434,6 +434,123 @@ async def test_typecheck_flag_value_general_error():
assert err.error_message == "Unknown flag type"


@pytest.mark.parametrize(
"flag_type, flag_value, default_value",
[
("integer", True, 1),
("integer", False, 0),
("float", True, 1.0),
("float", False, 0.0),
("boolean", 1, False),
("integer", "1", 0),
("float", 1, 0.0),
("string", True, "fallback"),
("object", True, {}),
],
)
@pytest.mark.parametrize("is_async", [False, True], ids=["sync", "async"])
@pytest.mark.asyncio
async def test_client_returns_default_on_type_mismatch(
flag_type, flag_value, default_value, is_async
):
provider = InMemoryProvider(
{"flag": InMemoryFlag("enabled", {"enabled": flag_value})}
)
api.set_provider_and_wait(provider)
client = get_client()
spy_hook = MagicMock(spec=Hook)
client.add_hooks([spy_hook])
suffix = "_async" if is_async else ""

get_details = getattr(client, f"get_{flag_type}_details{suffix}")
details = get_details("flag", default_value)
if is_async:
details = await details

assert details.flag_key == "flag"
assert details.value == default_value
# Equality alone cannot distinguish True from 1 (or False from 0).
assert type(details.value) is type(default_value)
assert details.reason == Reason.ERROR
assert details.error_code == ErrorCode.TYPE_MISMATCH
expected_type = (dict, list) if flag_type == "object" else type(default_value)
assert (
details.error_message
== f"Expected type {expected_type} but got {type(flag_value)}"
)
spy_hook.error.assert_called_once()
assert (
spy_hook.error.call_args.kwargs["exception"].error_code
== ErrorCode.TYPE_MISMATCH
)
spy_hook.after.assert_not_called()
spy_hook.finally_after.assert_called_once()
assert spy_hook.finally_after.call_args.kwargs["details"].value == default_value
assert type(spy_hook.finally_after.call_args.kwargs["details"].value) is type(
default_value
)

get_value = getattr(client, f"get_{flag_type}_value{suffix}")
value = get_value("flag", default_value)
if is_async:
value = await value
assert value == default_value
assert type(value) is type(default_value)


@pytest.mark.parametrize(
"flag_type, flag_value, default_value",
[
("boolean", True, False),
("boolean", False, True),
("integer", 0, -1),
("integer", 1, -1),
("integer", -1, 0),
("float", 0.0, -1.0),
("float", 1.5, -1.0),
("string", "enabled", "default"),
("object", {"enabled": True}, {}),
("object", [True, 1, "enabled"], []),
],
)
@pytest.mark.parametrize("is_async", [False, True], ids=["sync", "async"])
@pytest.mark.asyncio
async def test_client_preserves_matching_flag_types(
flag_type, flag_value, default_value, is_async
):
provider = InMemoryProvider(
{"flag": InMemoryFlag("enabled", {"enabled": flag_value})}
)
api.set_provider_and_wait(provider)
client = get_client()
suffix = "_async" if is_async else ""

get_details = getattr(client, f"get_{flag_type}_details{suffix}")
details = get_details("flag", default_value)
if is_async:
details = await details

assert details.value == flag_value
assert type(details.value) is type(flag_value)
assert details.variant == "enabled"
assert details.reason == Reason.STATIC
assert details.error_code is None

get_value = getattr(client, f"get_{flag_type}_value{suffix}")
value = get_value("flag", default_value)
if is_async:
value = await value
assert value == flag_value
assert type(value) is type(flag_value)


def test_typecheck_flag_value_accepts_integer_subclasses():
class IntegerValue(int):
pass

assert _typecheck_flag_value(IntegerValue(1), FlagType.INTEGER) is None


@pytest.mark.asyncio
async def test_typecheck_flag_value_type_mismatch_error():
# Given
Expand Down