From 6930816f209397d0323318179778de251bdadffb Mon Sep 17 00:00:00 2001 From: Ritinpaul Date: Tue, 6 Oct 2026 18:33:30 +0530 Subject: [PATCH] fix(hook): propagate evaluation context to subsequent before hooks (#628) Signed-off-by: Ritinpaul --- openfeature/hook/_hook_support.py | 50 ++- tests/hook/test_hook_support.py | 219 ++++++++++++ tests/test_before_hook_context_propagation.py | 318 ++++++++++++++++++ 3 files changed, 574 insertions(+), 13 deletions(-) create mode 100644 tests/test_before_hook_context_propagation.py diff --git a/openfeature/hook/_hook_support.py b/openfeature/hook/_hook_support.py index 4f292997..37ca5f21 100644 --- a/openfeature/hook/_hook_support.py +++ b/openfeature/hook/_hook_support.py @@ -1,6 +1,5 @@ import logging import typing -from functools import reduce from openfeature.evaluation_context import EvaluationContext from openfeature.flag_evaluation import FlagEvaluationDetails, FlagType @@ -59,19 +58,44 @@ def before_hooks( hooks_and_context: list[tuple[Hook, HookContext]], hints: HookHints | None = None, ) -> EvaluationContext: - kwargs = {"hints": hints} - executed_hooks = _execute_hooks_unchecked( - flag_type=flag_type, - hooks_and_context=hooks_and_context, - hook_method=HookType.BEFORE, - **kwargs, - ) - filtered_hooks = [result for result in executed_hooks if result is not None] - - if filtered_hooks: - return reduce(lambda a, b: a.merge(b), filtered_hooks) + # Requirement 4.3.4: Any evaluation context returned from a before hook MUST be + # passed to subsequent before hooks (via HookContext). + accumulated: EvaluationContext | None = None + supported_hooks_and_context = [ + (hook, hook_context) + for (hook, hook_context) in hooks_and_context + if hook.supports_flag_value_type(flag_type) + ] - return EvaluationContext() + try: + for hook, hook_context in supported_hooks_and_context: + if accumulated is not None: + # Propagate the accumulated context into this hook's HookContext so that + # it can observe the evaluation context returned by earlier before hooks. + if isinstance(hook_context.evaluation_context, EvaluationContext): + hook_context.evaluation_context = ( + hook_context.evaluation_context.merge(accumulated) + ) + else: + hook_context.evaluation_context = accumulated + + result = hook.before(hook_context=hook_context, hints=hints or {}) + + if isinstance(result, EvaluationContext): + accumulated = ( + accumulated.merge(result) if accumulated is not None else result + ) + finally: + if accumulated is not None: + for _, hook_context in supported_hooks_and_context: + if isinstance(hook_context.evaluation_context, EvaluationContext): + hook_context.evaluation_context = ( + hook_context.evaluation_context.merge(accumulated) + ) + else: + hook_context.evaluation_context = accumulated + + return accumulated if accumulated is not None else EvaluationContext() def _execute_hooks( diff --git a/tests/hook/test_hook_support.py b/tests/hook/test_hook_support.py index a72b0178..f8207eb2 100644 --- a/tests/hook/test_hook_support.py +++ b/tests/hook/test_hook_support.py @@ -122,6 +122,225 @@ def test_before_hooks_merges_evaluation_contexts(): assert context == EvaluationContext("bar", {"key_1": "val_1", "key_2": "val_2"}) +def test_before_hooks_propagates_context_to_subsequent_hook(): + # Given + initial_context = EvaluationContext(attributes={"initial": "present"}) + hook_context_a = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + hook_context_b = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + + received_by_hook_b: list[EvaluationContext] = [] + + hook_a = MagicMock(spec=Hook) + hook_a.before.return_value = EvaluationContext( + attributes={"from_hook_a": "visible"} + ) + + hook_b = MagicMock(spec=Hook) + + def hook_b_before(hook_context, hints): + received_by_hook_b.append(hook_context.evaluation_context) + return None + + hook_b.before.side_effect = hook_b_before + + # When + before_hooks( + FlagType.BOOLEAN, + [(hook_a, hook_context_a), (hook_b, hook_context_b)], + ) + + # Then + assert len(received_by_hook_b) == 1 + assert received_by_hook_b[0].attributes.get("from_hook_a") == "visible", ( + "Hook B did not receive the evaluation context returned by Hook A" + ) + assert received_by_hook_b[0].attributes.get("initial") == "present" + + +def test_before_hooks_accumulates_context_across_three_hooks(): + # Given + initial_context = EvaluationContext(attributes={"initial": "present"}) + ctx_a = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + ctx_b = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + ctx_c = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + + received_by_hook_b: list[EvaluationContext] = [] + received_by_hook_c: list[EvaluationContext] = [] + + hook_a = MagicMock(spec=Hook) + hook_a.before.return_value = EvaluationContext(attributes={"from_a": "A"}) + + hook_b = MagicMock(spec=Hook) + + def hook_b_before(hook_context, hints): + received_by_hook_b.append(hook_context.evaluation_context) + return EvaluationContext(attributes={"from_b": "B"}) + + hook_b.before.side_effect = hook_b_before + + hook_c = MagicMock(spec=Hook) + + def hook_c_before(hook_context, hints): + received_by_hook_c.append(hook_context.evaluation_context) + return None + + hook_c.before.side_effect = hook_c_before + + # When + before_hooks( + FlagType.BOOLEAN, + [(hook_a, ctx_a), (hook_b, ctx_b), (hook_c, ctx_c)], + ) + + # Then + assert received_by_hook_b[0].attributes.get("from_a") == "A", ( + "Hook B did not receive the evaluation context returned by Hook A" + ) + assert received_by_hook_c[0].attributes.get("from_a") == "A", ( + "Hook C did not receive the evaluation context returned by Hook A" + ) + assert received_by_hook_c[0].attributes.get("from_b") == "B", ( + "Hook C did not receive the evaluation context returned by Hook B" + ) + + +def test_before_hooks_later_hook_overrides_earlier_on_conflict(): + # Given + initial_context = EvaluationContext(attributes={"initial": "present"}) + ctx_a = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + ctx_b = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + ctx_c = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + + received_by_hook_c: list[EvaluationContext] = [] + + hook_a = MagicMock(spec=Hook) + hook_a.before.return_value = EvaluationContext(attributes={"shared": "A"}) + + hook_b = MagicMock(spec=Hook) + hook_b.before.return_value = EvaluationContext(attributes={"shared": "B"}) + + hook_c = MagicMock(spec=Hook) + + def hook_c_before(hook_context, hints): + received_by_hook_c.append(hook_context.evaluation_context) + return None + + hook_c.before.side_effect = hook_c_before + + # When + before_hooks( + FlagType.BOOLEAN, + [(hook_a, ctx_a), (hook_b, ctx_b), (hook_c, ctx_c)], + ) + + # Then + assert received_by_hook_c[0].attributes.get("shared") == "B", ( + "Later hook (B) result should override earlier hook (A) result for the same attribute" + ) + + +def test_before_hooks_none_result_does_not_corrupt_accumulated_context(): + # Given + initial_context = EvaluationContext(attributes={"initial": "present"}) + ctx_a = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + ctx_b = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + ctx_c = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + + received_by_hook_c: list[EvaluationContext] = [] + + hook_a = MagicMock(spec=Hook) + hook_a.before.return_value = EvaluationContext(attributes={"from_a": "A"}) + + hook_b = MagicMock(spec=Hook) + hook_b.before.return_value = None + + hook_c = MagicMock(spec=Hook) + + def hook_c_before(hook_context, hints): + received_by_hook_c.append(hook_context.evaluation_context) + return None + + hook_c.before.side_effect = hook_c_before + + # When + before_hooks( + FlagType.BOOLEAN, + [(hook_a, ctx_a), (hook_b, ctx_b), (hook_c, ctx_c)], + ) + + # Then + assert received_by_hook_c[0].attributes.get("from_a") == "A", ( + "Hook C should still see Hook A's context even though Hook B returned None" + ) + + +def test_before_hooks_finalizes_hook_contexts_for_all_participating_hooks(): + # Given + initial_context = EvaluationContext(attributes={"initial": "present"}) + ctx_a = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + ctx_b = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + + hook_a = MagicMock(spec=Hook) + hook_a.before.return_value = EvaluationContext(attributes={"from_a": "A"}) + hook_b = MagicMock(spec=Hook) + hook_b.before.return_value = EvaluationContext(attributes={"from_b": "B"}) + + # When + before_hooks(FlagType.BOOLEAN, [(hook_a, ctx_a), (hook_b, ctx_b)]) + + # Then + expected = {"initial": "present", "from_a": "A", "from_b": "B"} + assert ctx_a.evaluation_context.attributes == expected + assert ctx_b.evaluation_context.attributes == expected + + +def test_before_hooks_finalizes_hook_contexts_on_exception(): + # Given + initial_context = EvaluationContext(attributes={"initial": "present"}) + ctx_a = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + ctx_b = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + ctx_c = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + + hook_a = MagicMock(spec=Hook) + hook_a.before.return_value = EvaluationContext(attributes={"from_a": "A"}) + hook_b = MagicMock(spec=Hook) + hook_b.before.side_effect = RuntimeError("hook_b error") + hook_c = MagicMock(spec=Hook) + + # When + with pytest.raises(RuntimeError, match="hook_b error"): + before_hooks( + FlagType.BOOLEAN, [(hook_a, ctx_a), (hook_b, ctx_b), (hook_c, ctx_c)] + ) + + # Then + expected = {"initial": "present", "from_a": "A"} + assert hook_c.before.call_count == 0 + assert ctx_a.evaluation_context.attributes == expected + assert ctx_b.evaluation_context.attributes == expected + assert ctx_c.evaluation_context.attributes == expected + + +def test_before_hooks_unsupported_hook_context_is_not_finalized(): + # Given + initial_context = EvaluationContext(attributes={"initial": "present"}) + ctx_a = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + ctx_b = HookContext("flag_key", FlagType.BOOLEAN, True, initial_context) + + hook_a = MagicMock(spec=Hook) + hook_a.before.return_value = EvaluationContext(attributes={"from_a": "A"}) + hook_b = MagicMock(spec=Hook) + hook_b.supports_flag_value_type.return_value = False + + # When + before_hooks(FlagType.BOOLEAN, [(hook_a, ctx_a), (hook_b, ctx_b)]) + + # Then + assert hook_b.before.call_count == 0 + assert ctx_a.evaluation_context.attributes.get("from_a") == "A" + assert ctx_b.evaluation_context.attributes.get("from_a") is None + + def test_after_hooks_run_after_method(mock_hook): # Given hook_context = HookContext("flag_key", FlagType.BOOLEAN, True, "") diff --git a/tests/test_before_hook_context_propagation.py b/tests/test_before_hook_context_propagation.py new file mode 100644 index 00000000..58f2c96d --- /dev/null +++ b/tests/test_before_hook_context_propagation.py @@ -0,0 +1,318 @@ +import pytest + +from openfeature.api import get_client, set_provider_and_wait +from openfeature.evaluation_context import EvaluationContext +from openfeature.exception import GeneralError +from openfeature.flag_evaluation import Reason +from openfeature.hook import Hook +from openfeature.provider.no_op_provider import NoOpProvider + + +class CapturingHook(Hook): + def __init__(self, return_value=None, raise_in_before=False): + self.before_context: EvaluationContext | None = None + self.after_context: EvaluationContext | None = None + self.error_context: EvaluationContext | None = None + self.finally_context: EvaluationContext | None = None + self._return_value = return_value + self._raise_in_before = raise_in_before + + def before(self, hook_context, hints): + self.before_context = hook_context.evaluation_context + if self._raise_in_before: + raise RuntimeError("before hook failure") + return self._return_value + + def after(self, hook_context, details, hints): + self.after_context = hook_context.evaluation_context + + def error(self, hook_context, exception, hints): + self.error_context = hook_context.evaluation_context + + def finally_after(self, hook_context, details, hints): + self.finally_context = hook_context.evaluation_context + + @property + def received_context(self) -> EvaluationContext | None: + return self.before_context + + +class UnsupportedHook(CapturingHook): + def supports_flag_value_type(self, flag_type): + return False + + +class ContextCapturingProvider(NoOpProvider): + def __init__(self): + self.received_context: EvaluationContext | None = None + + def resolve_boolean_details(self, flag_key, default_value, evaluation_context): + self.received_context = evaluation_context + return super().resolve_boolean_details( + flag_key, default_value, evaluation_context + ) + + +class FailingProvider(NoOpProvider): + def resolve_boolean_details(self, flag_key, default_value, evaluation_context): + raise GeneralError("provider resolution failure") + + +def test_sync_second_before_hook_receives_context_from_first(): + # Given + set_provider_and_wait(NoOpProvider()) + client = get_client() + + hook_a = CapturingHook( + return_value=EvaluationContext(attributes={"from_hook_a": "visible"}) + ) + hook_b = CapturingHook(return_value=None) + client.add_hooks([hook_a, hook_b]) + + # When + client.get_boolean_value(flag_key="test-flag", default_value=False) + + # Then + assert hook_b.received_context is not None + assert hook_b.received_context.attributes.get("from_hook_a") == "visible", ( + "Hook B did not receive the evaluation context returned by Hook A (sync path)" + ) + + +def test_sync_third_before_hook_receives_accumulated_context(): + # Given + set_provider_and_wait(NoOpProvider()) + client = get_client() + + hook_a = CapturingHook(return_value=EvaluationContext(attributes={"from_a": "A"})) + hook_b = CapturingHook(return_value=EvaluationContext(attributes={"from_b": "B"})) + hook_c = CapturingHook(return_value=None) + client.add_hooks([hook_a, hook_b, hook_c]) + + # When + client.get_boolean_value(flag_key="test-flag", default_value=False) + + # Then + assert hook_b.received_context.attributes.get("from_a") == "A", ( + "Hook B did not receive the evaluation context returned by Hook A (sync path)" + ) + assert hook_c.received_context.attributes.get("from_a") == "A", ( + "Hook C did not receive Hook A's context (sync path)" + ) + assert hook_c.received_context.attributes.get("from_b") == "B", ( + "Hook C did not receive Hook B's context (sync path)" + ) + + +@pytest.mark.asyncio +async def test_async_second_before_hook_receives_context_from_first(): + # Given + set_provider_and_wait(NoOpProvider()) + client = get_client() + + hook_a = CapturingHook( + return_value=EvaluationContext(attributes={"from_hook_a": "visible"}) + ) + hook_b = CapturingHook(return_value=None) + client.add_hooks([hook_a, hook_b]) + + # When + await client.get_boolean_value_async(flag_key="test-flag", default_value=False) + + # Then + assert hook_b.received_context is not None + assert hook_b.received_context.attributes.get("from_hook_a") == "visible", ( + "Hook B did not receive the evaluation context returned by Hook A (async path)" + ) + + +@pytest.mark.asyncio +async def test_async_third_before_hook_receives_accumulated_context(): + # Given + set_provider_and_wait(NoOpProvider()) + client = get_client() + + hook_a = CapturingHook(return_value=EvaluationContext(attributes={"from_a": "A"})) + hook_b = CapturingHook(return_value=EvaluationContext(attributes={"from_b": "B"})) + hook_c = CapturingHook(return_value=None) + client.add_hooks([hook_a, hook_b, hook_c]) + + # When + await client.get_boolean_value_async(flag_key="test-flag", default_value=False) + + # Then + assert hook_b.received_context.attributes.get("from_a") == "A", ( + "Hook B did not receive the evaluation context returned by Hook A (async path)" + ) + assert hook_c.received_context.attributes.get("from_a") == "A", ( + "Hook C did not receive Hook A's context (async path)" + ) + assert hook_c.received_context.attributes.get("from_b") == "B", ( + "Hook C did not receive Hook B's context (async path)" + ) + + +def test_lifecycle_hooks_receive_accumulated_context_on_success(): + # Given + provider = ContextCapturingProvider() + set_provider_and_wait(provider) + client = get_client() + + hook_a = CapturingHook(return_value=EvaluationContext(attributes={"from_a": "A"})) + hook_b = CapturingHook(return_value=EvaluationContext(attributes={"from_b": "B"})) + client.add_hooks([hook_a, hook_b]) + + # When + client.get_boolean_value( + flag_key="test-flag", + default_value=False, + evaluation_context=EvaluationContext(attributes={"initial": "present"}), + ) + + # Then + expected = {"initial": "present", "from_a": "A", "from_b": "B"} + assert provider.received_context is not None + assert provider.received_context.attributes == expected + for hook in (hook_a, hook_b): + assert hook.after_context is not None + assert hook.after_context.attributes == expected + assert hook.finally_context is not None + assert hook.finally_context.attributes == expected + + +def test_lifecycle_hooks_receive_accumulated_context_on_before_failure(): + # Given + set_provider_and_wait(NoOpProvider()) + client = get_client() + + hook_a = CapturingHook(return_value=EvaluationContext(attributes={"from_a": "A"})) + hook_b = CapturingHook(raise_in_before=True) + hook_c = CapturingHook(return_value=None) + client.add_hooks([hook_a, hook_b, hook_c]) + + # When + res = client.get_boolean_value( + flag_key="test-flag", + default_value=False, + evaluation_context=EvaluationContext(attributes={"initial": "present"}), + ) + + # Then + expected = {"initial": "present", "from_a": "A"} + assert res is False + assert hook_c.before_context is None + for hook in (hook_a, hook_b, hook_c): + assert hook.error_context is not None + assert hook.error_context.attributes == expected + assert hook.finally_context is not None + assert hook.finally_context.attributes == expected + + +def test_lifecycle_hooks_receive_accumulated_context_on_provider_failure(): + # Given + set_provider_and_wait(FailingProvider()) + client = get_client() + + hook_a = CapturingHook(return_value=EvaluationContext(attributes={"from_a": "A"})) + hook_b = CapturingHook(return_value=EvaluationContext(attributes={"from_b": "B"})) + client.add_hooks([hook_a, hook_b]) + + # When + details = client.get_boolean_details( + flag_key="test-flag", + default_value=False, + evaluation_context=EvaluationContext(attributes={"initial": "present"}), + ) + + # Then + expected = {"initial": "present", "from_a": "A", "from_b": "B"} + assert details.reason == Reason.ERROR + for hook in (hook_a, hook_b): + assert hook.error_context is not None + assert hook.error_context.attributes == expected + assert hook.finally_context is not None + assert hook.finally_context.attributes == expected + + +def test_lifecycle_hooks_preserve_merge_precedence_on_conflict(): + # Given + provider = ContextCapturingProvider() + set_provider_and_wait(provider) + client = get_client() + + hook_a = CapturingHook(return_value=EvaluationContext(attributes={"shared": "A"})) + hook_b = CapturingHook(return_value=EvaluationContext(attributes={"shared": "B"})) + client.add_hooks([hook_a, hook_b]) + + # When + client.get_boolean_value( + flag_key="test-flag", + default_value=False, + evaluation_context=EvaluationContext(attributes={"shared": "initial"}), + ) + + # Then + assert hook_b.before_context is not None + assert hook_b.before_context.attributes.get("shared") == "A" + assert provider.received_context is not None + assert provider.received_context.attributes.get("shared") == "B" + for hook in (hook_a, hook_b): + assert hook.after_context is not None + assert hook.after_context.attributes.get("shared") == "B" + assert hook.finally_context is not None + assert hook.finally_context.attributes.get("shared") == "B" + + +def test_lifecycle_hooks_none_return_preserves_accumulated_context(): + # Given + provider = ContextCapturingProvider() + set_provider_and_wait(provider) + client = get_client() + + hook_a = CapturingHook(return_value=EvaluationContext(attributes={"from_a": "A"})) + hook_b = CapturingHook(return_value=None) + hook_c = CapturingHook(return_value=EvaluationContext(attributes={"from_c": "C"})) + client.add_hooks([hook_a, hook_b, hook_c]) + + # When + client.get_boolean_value( + flag_key="test-flag", + default_value=False, + evaluation_context=EvaluationContext(attributes={"initial": "present"}), + ) + + # Then + expected = {"initial": "present", "from_a": "A", "from_c": "C"} + assert hook_c.before_context is not None + assert hook_c.before_context.attributes.get("from_a") == "A" + assert provider.received_context is not None + assert provider.received_context.attributes == expected + for hook in (hook_a, hook_b, hook_c): + assert hook.after_context is not None + assert hook.after_context.attributes == expected + assert hook.finally_context is not None + assert hook.finally_context.attributes == expected + + +def test_lifecycle_hooks_unsupported_flag_type_context_is_not_finalized(): + # Given + set_provider_and_wait(NoOpProvider()) + client = get_client() + + hook_a = CapturingHook(return_value=EvaluationContext(attributes={"from_a": "A"})) + hook_b = UnsupportedHook() + client.add_hooks([hook_a, hook_b]) + + # When + client.get_boolean_value( + flag_key="test-flag", + default_value=False, + evaluation_context=EvaluationContext(attributes={"initial": "present"}), + ) + + # Then + assert hook_b.before_context is None + assert hook_b.after_context is None + assert hook_b.finally_context is None + assert hook_a.after_context is not None + assert hook_a.after_context.attributes.get("from_a") == "A"