From 9dcd7152cdfa1e38121b3373f78d300b63863069 Mon Sep 17 00:00:00 2001 From: tuhinkanti Date: Wed, 22 Jul 2026 15:01:59 -0700 Subject: [PATCH 1/3] feat(flagd): support custom gRPC metadata on in-process SyncFlags Add a `sync_metadata` option to FlagdProvider / Config that appends user-supplied gRPC metadata headers to every in-process SyncFlags call, alongside the provider-managed flagd-selector header. This lets callers inject infrastructure-specific headers on the long-lived sync stream -- e.g. `x-envoy-upstream-rq-timeout-ms: 0` to disable a proxy request timeout that would otherwise sever the stream after its default deadline. Keep `fatal_status_codes` as the last positional parameter in both public constructors so existing positional callers are unaffected, and add a regression test covering that ordering. Signed-off-by: Tuhin Sharma --- .../contrib/provider/flagd/config.py | 9 +++++ .../contrib/provider/flagd/provider.py | 5 +++ .../process/connector/grpc_watcher.py | 12 ++++-- .../tests/test_config.py | 18 +++++++++ .../tests/test_grpc_watcher.py | 39 +++++++++++++++++++ 5 files changed, 79 insertions(+), 4 deletions(-) diff --git a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/config.py b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/config.py index 75a2cb4b..d1beb694 100644 --- a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/config.py +++ b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/config.py @@ -105,6 +105,7 @@ def __init__( # noqa: PLR0913, PLR0915 channel_credentials: grpc.ChannelCredentials | None = None, sync_metadata_disabled: bool | None = None, fatal_status_codes: list[str] | None = None, + sync_metadata: typing.Sequence[tuple[str, str]] | None = None, ): self.host = env_or_default(ENV_VAR_HOST, DEFAULT_HOST) if host is None else host @@ -278,3 +279,11 @@ def __init__( # noqa: PLR0913, PLR0915 # Disabling will prevent static context from flagd being used in evaluations. # GetMetadata and this option will be removed. self.sync_metadata_disabled = sync_metadata_disabled + + # Additional gRPC metadata headers sent on every in-process SyncFlags call. + # Useful for injecting infrastructure-specific headers, e.g. disabling a + # proxy/mesh request timeout on the long-lived sync stream. These are merged + # with any headers the provider sets itself (such as ``flagd-selector``). + self.sync_metadata: tuple[tuple[str, str], ...] = ( + tuple(sync_metadata) if sync_metadata is not None else () + ) diff --git a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py index 0cbadbf7..ce5e4218 100644 --- a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py +++ b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py @@ -66,6 +66,7 @@ def __init__( # noqa: PLR0913 channel_credentials: grpc.ChannelCredentials | None = None, sync_metadata_disabled: bool | None = None, fatal_status_codes: list[str] | None = None, + sync_metadata: typing.Sequence[tuple[str, str]] | None = None, ): """ Create an instance of the FlagdProvider @@ -83,6 +84,9 @@ def __init__( # noqa: PLR0913 :param stream_deadline_ms: the maximum time to wait before a request times out :param keep_alive_time: the number of milliseconds to keep alive :param resolver_type: the type of resolver to use + :param sync_metadata: additional gRPC metadata headers (key-value tuples) sent on + every in-process SyncFlags call, e.g. to disable a proxy/mesh + request timeout on the long-lived sync stream """ if deadline_ms is None and timeout is not None: deadline_ms = timeout * 1000 @@ -112,6 +116,7 @@ def __init__( # noqa: PLR0913 default_authority=default_authority, channel_credentials=channel_credentials, sync_metadata_disabled=sync_metadata_disabled, + sync_metadata=sync_metadata, fatal_status_codes=fatal_status_codes, ) self.enriched_context: dict = {} diff --git a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py index e625ae33..8bde9217 100644 --- a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py +++ b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py @@ -212,17 +212,21 @@ def _create_request_args(self) -> dict: return request_args - def _create_metadata(self) -> tuple[tuple[str, str]] | None: + def _create_metadata(self) -> tuple[tuple[str, str], ...] | None: """Create gRPC metadata headers for the request. Returns gRPC metadata as a tuples of tuples containing header key-value pairs. The selector is passed via the 'flagd-selector' header per flagd v0.11.0+ specification, while also being included in the request body for backward compatibility with older flagd versions. + Any user-configured ``sync_metadata`` headers are appended, allowing callers to inject + infrastructure-specific headers (e.g. proxy/mesh timeout overrides) on the sync stream. """ - if self.selector is None: - return None + metadata: tuple[tuple[str, str], ...] = () + if self.selector is not None: + metadata += (("flagd-selector", self.selector),) + metadata += self.config.sync_metadata - return (("flagd-selector", self.selector),) + return metadata if metadata else None def _fetch_metadata(self) -> sync_pb2.GetMetadataResponse | None: if self.config.sync_metadata_disabled: diff --git a/providers/openfeature-provider-flagd/tests/test_config.py b/providers/openfeature-provider-flagd/tests/test_config.py index a671f8c1..e2a57d3b 100644 --- a/providers/openfeature-provider-flagd/tests/test_config.py +++ b/providers/openfeature-provider-flagd/tests/test_config.py @@ -44,6 +44,24 @@ def test_return_default_values_rpc(): assert config.retry_backoff_ms == DEFAULT_RETRY_BACKOFF assert config.stream_deadline_ms == DEFAULT_STREAM_DEADLINE assert config.tls is DEFAULT_TLS + assert config.sync_metadata == () + + +def test_sync_metadata_passthrough(): + metadata = [("x-envoy-upstream-rq-timeout-ms", "0")] + config = Config(resolver=ResolverType.IN_PROCESS, sync_metadata=metadata) + assert config.sync_metadata == (("x-envoy-upstream-rq-timeout-ms", "0"),) + + +def test_positional_fatal_status_codes_backwards_compatible(): + # fatal_status_codes must stay the last positional parameter so callers that + # passed it positionally before sync_metadata was added keep working. + # It is the 22nd positional parameter (21 parameters precede it). + leading_args = [None] * 21 + config = Config(*leading_args, ["UNAVAILABLE", "DATA_LOSS"]) + assert config.fatal_status_codes == ["UNAVAILABLE", "DATA_LOSS"] + # The positional value must not leak into sync_metadata. + assert config.sync_metadata == () def test_return_default_values_in_process(): diff --git a/providers/openfeature-provider-flagd/tests/test_grpc_watcher.py b/providers/openfeature-provider-flagd/tests/test_grpc_watcher.py index 395dd6a2..8e855366 100644 --- a/providers/openfeature-provider-flagd/tests/test_grpc_watcher.py +++ b/providers/openfeature-provider-flagd/tests/test_grpc_watcher.py @@ -36,6 +36,7 @@ def setUp(self): config.host = "localhost" config.port = 5000 config.sync_metadata_disabled = False + config.sync_metadata = () flag_store = Mock(spec=FlagStore) flag_store.update.return_value = None @@ -160,3 +161,41 @@ def test_selector_passed_via_both_metadata_and_body(self): self.assertIn("metadata", kwargs) metadata = kwargs["metadata"] self.assertEqual(metadata, (("flagd-selector", "test-selector"),)) + + def test_custom_sync_metadata_appended(self): + """User-configured sync_metadata headers are sent on the SyncFlags call.""" + self.grpc_watcher.selector = "test-selector" + self.grpc_watcher.config.sync_metadata = ( + ("x-envoy-upstream-rq-timeout-ms", "0"), + ) + mock_stream = iter( + [SyncFlagsResponse(flag_configuration='{"flag_key": "flag_value"}')] + ) + self.mock_stub.SyncFlags = Mock(return_value=mock_stream) + + self.run_listen_and_shutdown_after() + + metadata = self.mock_stub.SyncFlags.call_args.kwargs["metadata"] + self.assertEqual( + metadata, + ( + ("flagd-selector", "test-selector"), + ("x-envoy-upstream-rq-timeout-ms", "0"), + ), + ) + + def test_custom_sync_metadata_without_selector(self): + """sync_metadata is sent even when no selector is configured.""" + self.grpc_watcher.selector = None + self.grpc_watcher.config.sync_metadata = ( + ("x-envoy-upstream-rq-timeout-ms", "0"), + ) + mock_stream = iter( + [SyncFlagsResponse(flag_configuration='{"flag_key": "flag_value"}')] + ) + self.mock_stub.SyncFlags = Mock(return_value=mock_stream) + + self.run_listen_and_shutdown_after() + + metadata = self.mock_stub.SyncFlags.call_args.kwargs["metadata"] + self.assertEqual(metadata, (("x-envoy-upstream-rq-timeout-ms", "0"),)) From 81adc6e44dd0acbc58441946bb0e331ef8d23f44 Mon Sep 17 00:00:00 2001 From: Tuhin Sharma Date: Thu, 3 Sep 2026 15:20:37 -0700 Subject: [PATCH 2/3] feat(flagd): support gRPC client interceptors Signed-off-by: Tuhin Sharma --- .../openfeature-provider-flagd/README.md | 46 +++++++ .../contrib/provider/flagd/config.py | 29 ++++- .../contrib/provider/flagd/provider.py | 10 +- .../contrib/provider/flagd/resolvers/grpc.py | 9 +- .../process/connector/grpc_watcher.py | 25 ++-- .../tests/test_config.py | 21 +-- .../tests/test_grpc_resolver.py | 88 ++++++++++++- .../tests/test_grpc_watcher.py | 123 +++++++++++++----- 8 files changed, 279 insertions(+), 72 deletions(-) diff --git a/providers/openfeature-provider-flagd/README.md b/providers/openfeature-provider-flagd/README.md index fbade51f..845dd5fc 100644 --- a/providers/openfeature-provider-flagd/README.md +++ b/providers/openfeature-provider-flagd/README.md @@ -92,6 +92,9 @@ The default options can be defined in the FlagdProvider constructor. | max_cache_size | FLAGD_MAX_CACHE_SIZE | int | 1000 | rpc | | retry_backoff_ms | FLAGD_RETRY_BACKOFF_MS | int | 1000 | rpc | | offline_flag_source_path | FLAGD_OFFLINE_FLAG_SOURCE_PATH | str | null | in-process | +| sync_metadata_disabled | - | bool | null | in-process | +| fatal_status_codes | FLAGD_FATAL_STATUS_CODES | sequence of gRPC status code names | empty | rpc & in-process | +| client_interceptors | - | sequence of gRPC client interceptors | null | rpc & in-process | > [!NOTE] > The `selector` configuration is only used in **in-process** mode for filtering flag configurations. See [Selector Handling](#selector-handling-in-process-mode-only) for migration guidance. @@ -106,6 +109,49 @@ The default options can be defined in the FlagdProvider constructor. > [!NOTE] > Some configurations are only applicable for RPC resolver. +### Custom gRPC interceptors + +`client_interceptors` are synchronous gRPC client interceptors applied to the channel in the order provided. Use them for infrastructure concerns such as custom headers or credentials. Flagd-specific options like `selector` stay first-class and do not need a custom interceptor. + +Metadata keys added by an interceptor must be valid lowercase gRPC metadata keys. If an interceptor adds `flagd-selector` while `selector` is set, the request contains duplicate keys. gRPC permits duplicate metadata keys. + +`grpc.aio` interceptors are not supported. Passing an object that does not implement one of the synchronous client interceptor interfaces raises `TypeError` when the provider creates its channel. +Exceptions raised while opening a sync or event stream are logged and retried after `retry_backoff_max_ms`; a persistently failing interceptor prevents stream updates. + +```python +import grpc +from openfeature.contrib.provider.flagd import FlagdProvider +from openfeature.contrib.provider.flagd.config import ResolverType + + +class _ClientCallDetails(grpc.ClientCallDetails): + def __init__(self, details, metadata): + self.method = details.method + self.timeout = details.timeout + self.metadata = metadata + self.credentials = details.credentials + self.wait_for_ready = details.wait_for_ready + self.compression = details.compression + + +class DisableEnvoyTimeout(grpc.UnaryStreamClientInterceptor): + def intercept_unary_stream(self, continuation, client_call_details, request): + metadata = list(client_call_details.metadata or []) + metadata.append(("x-envoy-upstream-rq-timeout-ms", "0")) + details = _ClientCallDetails(client_call_details, metadata) + return continuation(details, request) + + +provider = FlagdProvider( + resolver_type=ResolverType.IN_PROCESS, + client_interceptors=[DisableEnvoyTimeout()], +) +``` + +See also: +- https://grpc.github.io/grpc/python/grpc.html#client-side-interceptor +- https://grpc.github.io/grpc/python/grpc.html#grpc.intercept_channel + ### Selector Handling (In-Process Mode Only) > [!IMPORTANT] diff --git a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/config.py b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/config.py index d1beb694..90dd7aee 100644 --- a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/config.py +++ b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/config.py @@ -57,6 +57,22 @@ class CacheType(Enum): T = typing.TypeVar("T") +ClientInterceptor: typing.TypeAlias = ( + grpc.UnaryUnaryClientInterceptor + | grpc.UnaryStreamClientInterceptor + | grpc.StreamUnaryClientInterceptor + | grpc.StreamStreamClientInterceptor +) + + +def apply_client_interceptors( + channel: grpc.Channel, + client_interceptors: typing.Sequence[ClientInterceptor], +) -> grpc.Channel: + if not client_interceptors: + return channel + return grpc.intercept_channel(channel, *client_interceptors) + def str_to_bool(val: str) -> bool: return val.lower() == "true" @@ -105,7 +121,7 @@ def __init__( # noqa: PLR0913, PLR0915 channel_credentials: grpc.ChannelCredentials | None = None, sync_metadata_disabled: bool | None = None, fatal_status_codes: list[str] | None = None, - sync_metadata: typing.Sequence[tuple[str, str]] | None = None, + client_interceptors: typing.Sequence[ClientInterceptor] | None = None, ): self.host = env_or_default(ENV_VAR_HOST, DEFAULT_HOST) if host is None else host @@ -280,10 +296,9 @@ def __init__( # noqa: PLR0913, PLR0915 # GetMetadata and this option will be removed. self.sync_metadata_disabled = sync_metadata_disabled - # Additional gRPC metadata headers sent on every in-process SyncFlags call. - # Useful for injecting infrastructure-specific headers, e.g. disabling a - # proxy/mesh request timeout on the long-lived sync stream. These are merged - # with any headers the provider sets itself (such as ``flagd-selector``). - self.sync_metadata: tuple[tuple[str, str], ...] = ( - tuple(sync_metadata) if sync_metadata is not None else () + # gRPC client interceptors applied to the channel (rpc and in-process). + # Use this for infrastructure concerns such as custom headers or + # credentials; flagd-specific options (e.g. selector) stay first-class. + self.client_interceptors: tuple[ClientInterceptor, ...] = ( + tuple(client_interceptors) if client_interceptors is not None else () ) diff --git a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py index ce5e4218..1d76cdd2 100644 --- a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py +++ b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py @@ -33,7 +33,7 @@ from openfeature.provider import AbstractProvider from openfeature.provider.metadata import Metadata -from .config import CacheType, Config, ResolverType +from .config import CacheType, ClientInterceptor, Config, ResolverType from .resolvers import AbstractResolver, GrpcResolver, InProcessResolver from .sync_metadata_hook import SyncMetadataHook @@ -66,7 +66,7 @@ def __init__( # noqa: PLR0913 channel_credentials: grpc.ChannelCredentials | None = None, sync_metadata_disabled: bool | None = None, fatal_status_codes: list[str] | None = None, - sync_metadata: typing.Sequence[tuple[str, str]] | None = None, + client_interceptors: typing.Sequence[ClientInterceptor] | None = None, ): """ Create an instance of the FlagdProvider @@ -84,9 +84,7 @@ def __init__( # noqa: PLR0913 :param stream_deadline_ms: the maximum time to wait before a request times out :param keep_alive_time: the number of milliseconds to keep alive :param resolver_type: the type of resolver to use - :param sync_metadata: additional gRPC metadata headers (key-value tuples) sent on - every in-process SyncFlags call, e.g. to disable a proxy/mesh - request timeout on the long-lived sync stream + :param client_interceptors: gRPC client interceptors applied to the channel. Metadata keys added by interceptors must be valid lowercase gRPC metadata keys. An interceptor that adds ``flagd-selector`` can duplicate the provider's selector metadata. """ if deadline_ms is None and timeout is not None: deadline_ms = timeout * 1000 @@ -116,8 +114,8 @@ def __init__( # noqa: PLR0913 default_authority=default_authority, channel_credentials=channel_credentials, sync_metadata_disabled=sync_metadata_disabled, - sync_metadata=sync_metadata, fatal_status_codes=fatal_status_codes, + client_interceptors=client_interceptors, ) self.enriched_context: dict = {} diff --git a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/grpc.py b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/grpc.py index 919fa259..8cc68cee 100644 --- a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/grpc.py +++ b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/grpc.py @@ -28,7 +28,7 @@ evaluation_pb2_grpc, ) -from ..config import CacheType, Config +from ..config import CacheType, Config, apply_client_interceptors from ..flag_type import FlagType from .types import GrpcMultiCallableArgs @@ -135,7 +135,7 @@ def _generate_channel(self, config: Config) -> grpc.Channel: options=options, ) - return channel + return apply_client_interceptors(channel, config.client_interceptors) def initialize(self, evaluation_context: EvaluationContext) -> None: self.connect() @@ -289,6 +289,11 @@ def listen(self) -> None: logger.exception( f"Could not parse flag data using flagd syntax: {message=}" ) + except Exception: + if self.active: + logger.exception("Unexpected EventStream error, reconnecting") + else: + logger.debug("EventStream ended during shutdown", exc_info=True) if self.active: self._wait_before_reconnect() diff --git a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py index 955a8cb6..29a5855c 100644 --- a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py +++ b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py @@ -21,7 +21,7 @@ sync_pb2_grpc, ) -from ....config import Config +from ....config import Config, apply_client_interceptors from ...types import GrpcMultiCallableArgs from ..connector import FlagStateConnector from ..flags import FlagStore @@ -122,7 +122,7 @@ def _generate_channel(self, config: Config) -> grpc.Channel: options=options, ) - return channel + return apply_client_interceptors(channel, config.client_interceptors) def initialize(self, context: EvaluationContext) -> None: self.connect() @@ -214,21 +214,17 @@ def _create_request_args(self) -> dict: return request_args - def _create_metadata(self) -> tuple[tuple[str, str], ...] | None: + def _create_metadata(self) -> tuple[tuple[str, str]] | None: """Create gRPC metadata headers for the request. Returns gRPC metadata as a tuples of tuples containing header key-value pairs. The selector is passed via the 'flagd-selector' header per flagd v0.11.0+ specification, while also being included in the request body for backward compatibility with older flagd versions. - Any user-configured ``sync_metadata`` headers are appended, allowing callers to inject - infrastructure-specific headers (e.g. proxy/mesh timeout overrides) on the sync stream. """ - metadata: tuple[tuple[str, str], ...] = () - if self.selector is not None: - metadata += (("flagd-selector", self.selector),) - metadata += self.config.sync_metadata + if self.selector is None: + return None - return metadata if metadata else None + return (("flagd-selector", self.selector),) def _fetch_metadata(self) -> sync_pb2.GetMetadataResponse | None: if self.config.sync_metadata_disabled: @@ -297,7 +293,7 @@ def _handle_rpc_error(self, e: grpc.RpcError) -> bool: def _wait_before_reconnect(self) -> None: self._shutdown_event.wait(self.retry_backoff_max_seconds) - def listen(self) -> None: + def listen(self) -> None: # noqa: C901 call_args = self.generate_grpc_call_args() request_args = self._create_request_args() @@ -318,6 +314,13 @@ def listen(self) -> None: ) except ParseError: logger.exception("Could not parse flag data using flagd syntax") + except Exception: + if self.active: + logger.exception("Unexpected SyncFlags stream error, reconnecting") + else: + logger.debug( + "SyncFlags stream ended during shutdown", exc_info=True + ) if self.active: self._wait_before_reconnect() diff --git a/providers/openfeature-provider-flagd/tests/test_config.py b/providers/openfeature-provider-flagd/tests/test_config.py index e2a57d3b..017b853f 100644 --- a/providers/openfeature-provider-flagd/tests/test_config.py +++ b/providers/openfeature-provider-flagd/tests/test_config.py @@ -1,3 +1,6 @@ +from unittest.mock import Mock + +import grpc import pytest # not sure if we still need this test, as this is also covered with gherkin tests. @@ -44,24 +47,22 @@ def test_return_default_values_rpc(): assert config.retry_backoff_ms == DEFAULT_RETRY_BACKOFF assert config.stream_deadline_ms == DEFAULT_STREAM_DEADLINE assert config.tls is DEFAULT_TLS - assert config.sync_metadata == () + assert config.client_interceptors == () -def test_sync_metadata_passthrough(): - metadata = [("x-envoy-upstream-rq-timeout-ms", "0")] - config = Config(resolver=ResolverType.IN_PROCESS, sync_metadata=metadata) - assert config.sync_metadata == (("x-envoy-upstream-rq-timeout-ms", "0"),) +def test_client_interceptors_passthrough(): + interceptor = Mock(spec=grpc.UnaryUnaryClientInterceptor) + config = Config(resolver=ResolverType.IN_PROCESS, client_interceptors=[interceptor]) + assert config.client_interceptors == (interceptor,) def test_positional_fatal_status_codes_backwards_compatible(): - # fatal_status_codes must stay the last positional parameter so callers that - # passed it positionally before sync_metadata was added keep working. - # It is the 22nd positional parameter (21 parameters precede it). + # fatal_status_codes stays ahead of client_interceptors so callers that + # passed it positionally keep working. It is the 22nd positional parameter. leading_args = [None] * 21 config = Config(*leading_args, ["UNAVAILABLE", "DATA_LOSS"]) assert config.fatal_status_codes == ["UNAVAILABLE", "DATA_LOSS"] - # The positional value must not leak into sync_metadata. - assert config.sync_metadata == () + assert config.client_interceptors == () def test_return_default_values_in_process(): diff --git a/providers/openfeature-provider-flagd/tests/test_grpc_resolver.py b/providers/openfeature-provider-flagd/tests/test_grpc_resolver.py index ccc8ee29..6f976fa0 100644 --- a/providers/openfeature-provider-flagd/tests/test_grpc_resolver.py +++ b/providers/openfeature-provider-flagd/tests/test_grpc_resolver.py @@ -1,5 +1,4 @@ import unittest -from contextlib import suppress from unittest.mock import MagicMock, Mock, patch import grpc @@ -66,12 +65,16 @@ def test_unary_call_omits_metadata_when_no_selector(): def test_event_stream_includes_selector_metadata_when_configured(): resolver = _make_resolver("test-selector") mock_stub = MagicMock() - mock_stub.EventStream = Mock(side_effect=Exception("break loop")) + + def stop_after_call(*args, **kwargs): + resolver.active = False + return iter(()) + + mock_stub.EventStream = Mock(side_effect=stop_after_call) resolver.stub = mock_stub resolver.active = True - with suppress(Exception): - resolver.listen() + resolver.listen() kwargs = mock_stub.EventStream.call_args.kwargs assert kwargs.get("metadata") == ((FLAGD_SELECTOR_HEADER, "test-selector"),) @@ -80,12 +83,16 @@ def test_event_stream_includes_selector_metadata_when_configured(): def test_event_stream_omits_metadata_when_no_selector(): resolver = _make_resolver(None) mock_stub = MagicMock() - mock_stub.EventStream = Mock(side_effect=Exception("break loop")) + + def stop_after_call(*args, **kwargs): + resolver.active = False + return iter(()) + + mock_stub.EventStream = Mock(side_effect=stop_after_call) resolver.stub = mock_stub resolver.active = True - with suppress(Exception): - resolver.listen() + resolver.listen() kwargs = mock_stub.EventStream.call_args.kwargs assert "metadata" not in kwargs @@ -144,6 +151,73 @@ def test_listen_backs_off_after_stream_completion(self): wait_before_reconnect.assert_called_once() + def test_listen_backs_off_after_unexpected_error(self): + self.grpc_resolver.stub.EventStream = Mock( + side_effect=RuntimeError("interceptor failed") + ) + + with patch.object( + self.grpc_resolver, + "_wait_before_reconnect", + side_effect=lambda: setattr(self.grpc_resolver, "active", False), + ) as wait_before_reconnect: + self.grpc_resolver.listen() + + wait_before_reconnect.assert_called_once() + + def test_generate_channel_applies_client_interceptors(self): + interceptor = Mock(spec=grpc.UnaryUnaryClientInterceptor) + raw_channel = Mock(spec=Channel) + wrapped_channel = Mock(spec=Channel) + config = Config( + tls=False, cache=CacheType.DISABLED, client_interceptors=[interceptor] + ) + + with ( + patch( + "openfeature.contrib.provider.flagd.resolvers.grpc.grpc.insecure_channel", + return_value=raw_channel, + ), + patch( + "openfeature.contrib.provider.flagd.config.grpc.intercept_channel", + return_value=wrapped_channel, + ) as intercept_channel, + ): + resolver = GrpcResolver( + config=config, + emit_provider_ready=Mock(), + emit_provider_error=Mock(), + emit_provider_stale=Mock(), + emit_provider_configuration_changed=Mock(), + ) + + self.assertIs(resolver.channel, wrapped_channel) + intercept_channel.assert_called_once_with(raw_channel, interceptor) + + def test_generate_channel_skips_intercept_channel_when_no_interceptors(self): + raw_channel = Mock(spec=Channel) + config = Config(tls=False, cache=CacheType.DISABLED) + + with ( + patch( + "openfeature.contrib.provider.flagd.resolvers.grpc.grpc.insecure_channel", + return_value=raw_channel, + ), + patch( + "openfeature.contrib.provider.flagd.config.grpc.intercept_channel", + ) as intercept_channel, + ): + resolver = GrpcResolver( + config=config, + emit_provider_ready=Mock(), + emit_provider_error=Mock(), + emit_provider_stale=Mock(), + emit_provider_configuration_changed=Mock(), + ) + + self.assertIs(resolver.channel, raw_channel) + intercept_channel.assert_not_called() + if __name__ == "__main__": unittest.main() diff --git a/providers/openfeature-provider-flagd/tests/test_grpc_watcher.py b/providers/openfeature-provider-flagd/tests/test_grpc_watcher.py index d5fb8762..33cc8820 100644 --- a/providers/openfeature-provider-flagd/tests/test_grpc_watcher.py +++ b/providers/openfeature-provider-flagd/tests/test_grpc_watcher.py @@ -16,6 +16,7 @@ from openfeature.event import ProviderEventDetails from openfeature.schemas.protobuf.flagd.sync.v1.sync_pb2 import ( GetMetadataResponse, + SyncFlagsRequest, SyncFlagsResponse, ) from openfeature.schemas.protobuf.flagd.sync.v1.sync_pb2_grpc import FlagSyncServiceStub @@ -29,6 +30,24 @@ def details(self): return "stream unavailable" +class _ClientCallDetails(grpc.ClientCallDetails): + def __init__(self, details, metadata): + self.method = details.method + self.timeout = details.timeout + self.metadata = metadata + self.credentials = details.credentials + self.wait_for_ready = details.wait_for_ready + self.compression = details.compression + + +class _MetadataInterceptor(grpc.UnaryStreamClientInterceptor): + def intercept_unary_stream(self, continuation, client_call_details, request): + metadata = list(client_call_details.metadata or []) + metadata.append(("x-envoy-upstream-rq-timeout-ms", "0")) + details = _ClientCallDetails(client_call_details, metadata) + return continuation(details, request) + + class TestGrpcWatcher(unittest.TestCase): def setUp(self): config = Mock(spec=Config) @@ -45,7 +64,6 @@ def setUp(self): config.host = "localhost" config.port = 5000 config.sync_metadata_disabled = False - config.sync_metadata = () config.fatal_status_codes = [] flag_store = Mock(spec=FlagStore) @@ -171,6 +189,18 @@ def test_listen_backs_off_after_stream_completion(self): wait_before_reconnect.assert_called_once() + def test_listen_backs_off_after_unexpected_error(self): + self.mock_stub.SyncFlags = Mock(side_effect=RuntimeError("interceptor failed")) + + with patch.object( + self.grpc_watcher, + "_wait_before_reconnect", + side_effect=lambda: setattr(self.grpc_watcher, "active", False), + ) as wait_before_reconnect: + self.grpc_watcher.listen() + + wait_before_reconnect.assert_called_once() + def test_selector_passed_via_both_metadata_and_body(self): """Test that selector is passed via both gRPC metadata header and request body for backward compatibility""" self.grpc_watcher.selector = "test-selector" @@ -199,40 +229,75 @@ def test_selector_passed_via_both_metadata_and_body(self): metadata = kwargs["metadata"] self.assertEqual(metadata, (("flagd-selector", "test-selector"),)) - def test_custom_sync_metadata_appended(self): - """User-configured sync_metadata headers are sent on the SyncFlags call.""" - self.grpc_watcher.selector = "test-selector" - self.grpc_watcher.config.sync_metadata = ( - ("x-envoy-upstream-rq-timeout-ms", "0"), - ) - mock_stream = iter( - [SyncFlagsResponse(flag_configuration='{"flag_key": "flag_value"}')] - ) - self.mock_stub.SyncFlags = Mock(return_value=mock_stream) + def test_client_interceptor_adds_metadata_to_sync_flags(self): + raw_channel = Mock(spec=Channel) + raw_sync_flags = Mock(return_value=iter(())) + raw_channel.unary_stream.return_value = raw_sync_flags + config = Config(tls=False, client_interceptors=[_MetadataInterceptor()]) - self.run_listen_and_shutdown_after() + with patch( + "openfeature.contrib.provider.flagd.resolvers.process.connector.grpc_watcher.grpc.insecure_channel", + return_value=raw_channel, + ): + watcher = GrpcWatcher( + config=config, + flag_store=Mock(spec=FlagStore), + emit_provider_ready=Mock(), + emit_provider_error=Mock(), + emit_provider_stale=Mock(), + ) + + watcher.stub.SyncFlags( + SyncFlagsRequest(), metadata=(("flagd-selector", "test-selector"),) + ) - metadata = self.mock_stub.SyncFlags.call_args.kwargs["metadata"] self.assertEqual( - metadata, - ( + raw_sync_flags.call_args.kwargs["metadata"], + [ ("flagd-selector", "test-selector"), ("x-envoy-upstream-rq-timeout-ms", "0"), - ), + ], ) - def test_custom_sync_metadata_without_selector(self): - """sync_metadata is sent even when no selector is configured.""" - self.grpc_watcher.selector = None - self.grpc_watcher.config.sync_metadata = ( - ("x-envoy-upstream-rq-timeout-ms", "0"), - ) - mock_stream = iter( - [SyncFlagsResponse(flag_configuration='{"flag_key": "flag_value"}')] - ) - self.mock_stub.SyncFlags = Mock(return_value=mock_stream) + def test_generate_channel_rejects_invalid_client_interceptor(self): + raw_channel = Mock(spec=Channel) + config = Config(tls=False, client_interceptors=[object()]) - self.run_listen_and_shutdown_after() + with ( + patch( + "openfeature.contrib.provider.flagd.resolvers.process.connector.grpc_watcher.grpc.insecure_channel", + return_value=raw_channel, + ), + self.assertRaises(TypeError), + ): + GrpcWatcher( + config=config, + flag_store=Mock(spec=FlagStore), + emit_provider_ready=Mock(), + emit_provider_error=Mock(), + emit_provider_stale=Mock(), + ) + + def test_generate_channel_skips_intercept_channel_when_no_interceptors(self): + raw_channel = Mock(spec=Channel) + config = Config(tls=False) + + with ( + patch( + "openfeature.contrib.provider.flagd.resolvers.process.connector.grpc_watcher.grpc.insecure_channel", + return_value=raw_channel, + ), + patch( + "openfeature.contrib.provider.flagd.config.grpc.intercept_channel", + ) as intercept_channel, + ): + watcher = GrpcWatcher( + config=config, + flag_store=Mock(spec=FlagStore), + emit_provider_ready=Mock(), + emit_provider_error=Mock(), + emit_provider_stale=Mock(), + ) - metadata = self.mock_stub.SyncFlags.call_args.kwargs["metadata"] - self.assertEqual(metadata, (("x-envoy-upstream-rq-timeout-ms", "0"),)) + self.assertIs(watcher.channel, raw_channel) + intercept_channel.assert_not_called() From ebbb279aa39b3ad5be04fd23cf6bd04b75492fb4 Mon Sep 17 00:00:00 2001 From: Tuhin Sharma Date: Thu, 3 Sep 2026 21:01:58 -0700 Subject: [PATCH 3/3] fix(flagd): support custom channel credentials Signed-off-by: Tuhin Sharma --- .../openfeature-provider-flagd/README.md | 19 +++++++ .../contrib/provider/flagd/provider.py | 1 + .../contrib/provider/flagd/resolvers/grpc.py | 9 +++- .../process/connector/grpc_watcher.py | 51 ++++++++++--------- .../tests/test_grpc_resolver.py | 39 ++++++++++++++ 5 files changed, 94 insertions(+), 25 deletions(-) diff --git a/providers/openfeature-provider-flagd/README.md b/providers/openfeature-provider-flagd/README.md index 845dd5fc..87c52878 100644 --- a/providers/openfeature-provider-flagd/README.md +++ b/providers/openfeature-provider-flagd/README.md @@ -94,6 +94,7 @@ The default options can be defined in the FlagdProvider constructor. | offline_flag_source_path | FLAGD_OFFLINE_FLAG_SOURCE_PATH | str | null | in-process | | sync_metadata_disabled | - | bool | null | in-process | | fatal_status_codes | FLAGD_FATAL_STATUS_CODES | sequence of gRPC status code names | empty | rpc & in-process | +| channel_credentials | - | `grpc.ChannelCredentials` (including mTLS) | null | rpc & in-process | | client_interceptors | - | sequence of gRPC client interceptors | null | rpc & in-process | > [!NOTE] @@ -109,6 +110,24 @@ The default options can be defined in the FlagdProvider constructor. > [!NOTE] > Some configurations are only applicable for RPC resolver. +### Mutual TLS + +Pass custom `grpc.ChannelCredentials` to `channel_credentials` when the flagd server requires mutual TLS (mTLS). The provider uses these credentials for both resolver types and gives them precedence over `tls` and `cert_path`. + +```python +import grpc + +from openfeature.contrib.provider.flagd import FlagdProvider + +credentials = grpc.ssl_channel_credentials( + root_certificates=ca_certificate, + private_key=client_private_key, + certificate_chain=client_certificate, +) + +provider = FlagdProvider(channel_credentials=credentials) +``` + ### Custom gRPC interceptors `client_interceptors` are synchronous gRPC client interceptors applied to the channel in the order provided. Use them for infrastructure concerns such as custom headers or credentials. Flagd-specific options like `selector` stay first-class and do not need a custom interceptor. diff --git a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py index 1d76cdd2..6d0fa000 100644 --- a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py +++ b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/provider.py @@ -84,6 +84,7 @@ def __init__( # noqa: PLR0913 :param stream_deadline_ms: the maximum time to wait before a request times out :param keep_alive_time: the number of milliseconds to keep alive :param resolver_type: the type of resolver to use + :param channel_credentials: custom gRPC channel credentials, including mTLS credentials :param client_interceptors: gRPC client interceptors applied to the channel. Metadata keys added by interceptors must be valid lowercase gRPC metadata keys. An interceptor that adds ``flagd-selector`` can duplicate the provider's selector metadata. """ if deadline_ms is None and timeout is not None: diff --git a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/grpc.py b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/grpc.py index 8cc68cee..b3274865 100644 --- a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/grpc.py +++ b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/grpc.py @@ -117,7 +117,14 @@ def _generate_channel(self, config: Config) -> grpc.Channel: ), ), ] - if config.tls: + if config.channel_credentials is not None: + channel = grpc.secure_channel( + target, + credentials=config.channel_credentials, + options=options, + ) + + elif config.tls: credentials = grpc.ssl_channel_credentials() if config.cert_path: with open(config.cert_path, "rb") as f: diff --git a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py index 29a5855c..fb50c69a 100644 --- a/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py +++ b/providers/openfeature-provider-flagd/src/openfeature/contrib/provider/flagd/resolvers/process/connector/grpc_watcher.py @@ -293,34 +293,37 @@ def _handle_rpc_error(self, e: grpc.RpcError) -> bool: def _wait_before_reconnect(self) -> None: self._shutdown_event.wait(self.retry_backoff_max_seconds) - def listen(self) -> None: # noqa: C901 + def _listen_once( + self, call_args: GrpcMultiCallableArgs, request_args: dict + ) -> bool: + try: + context_values_response = self._fetch_metadata() + request = sync_pb2.SyncFlagsRequest(**request_args) + logger.debug("Setting up gRPC sync flags connection") + for flag_rsp in self.stub.SyncFlags(request, **call_args): + if self._handle_flag_response(flag_rsp, context_values_response): + return True + except grpc.RpcError as e: + if self._handle_rpc_error(e): + return True + except json.JSONDecodeError: + logger.exception("Could not parse JSON flag data from SyncFlags endpoint") + except ParseError: + logger.exception("Could not parse flag data using flagd syntax") + except Exception: + if self.active: + logger.exception("Unexpected SyncFlags stream error, reconnecting") + else: + logger.debug("SyncFlags stream ended during shutdown", exc_info=True) + return False + + def listen(self) -> None: call_args = self.generate_grpc_call_args() request_args = self._create_request_args() while self.active: - try: - context_values_response = self._fetch_metadata() - request = sync_pb2.SyncFlagsRequest(**request_args) - logger.debug("Setting up gRPC sync flags connection") - for flag_rsp in self.stub.SyncFlags(request, **call_args): - if self._handle_flag_response(flag_rsp, context_values_response): - return - except grpc.RpcError as e: - if self._handle_rpc_error(e): - return - except json.JSONDecodeError: - logger.exception( - "Could not parse JSON flag data from SyncFlags endpoint" - ) - except ParseError: - logger.exception("Could not parse flag data using flagd syntax") - except Exception: - if self.active: - logger.exception("Unexpected SyncFlags stream error, reconnecting") - else: - logger.debug( - "SyncFlags stream ended during shutdown", exc_info=True - ) + if self._listen_once(call_args, request_args): + return if self.active: self._wait_before_reconnect() diff --git a/providers/openfeature-provider-flagd/tests/test_grpc_resolver.py b/providers/openfeature-provider-flagd/tests/test_grpc_resolver.py index 6f976fa0..f558757e 100644 --- a/providers/openfeature-provider-flagd/tests/test_grpc_resolver.py +++ b/providers/openfeature-provider-flagd/tests/test_grpc_resolver.py @@ -194,6 +194,45 @@ def test_generate_channel_applies_client_interceptors(self): self.assertIs(resolver.channel, wrapped_channel) intercept_channel.assert_called_once_with(raw_channel, interceptor) + def test_generate_channel_uses_custom_channel_credentials(self): + credentials = Mock(spec=grpc.ChannelCredentials) + interceptor = Mock(spec=grpc.UnaryUnaryClientInterceptor) + raw_channel = Mock(spec=Channel) + wrapped_channel = Mock(spec=Channel) + config = Config( + tls=False, + cache=CacheType.DISABLED, + channel_credentials=credentials, + client_interceptors=[interceptor], + ) + + with ( + patch( + "openfeature.contrib.provider.flagd.resolvers.grpc.grpc.secure_channel", + return_value=raw_channel, + ) as secure_channel, + patch( + "openfeature.contrib.provider.flagd.resolvers.grpc.grpc.insecure_channel", + ) as insecure_channel, + patch( + "openfeature.contrib.provider.flagd.config.grpc.intercept_channel", + return_value=wrapped_channel, + ) as intercept_channel, + ): + resolver = GrpcResolver( + config=config, + emit_provider_ready=Mock(), + emit_provider_error=Mock(), + emit_provider_stale=Mock(), + emit_provider_configuration_changed=Mock(), + ) + + self.assertIs(resolver.channel, wrapped_channel) + secure_channel.assert_called_once() + self.assertIs(secure_channel.call_args.kwargs["credentials"], credentials) + insecure_channel.assert_not_called() + intercept_channel.assert_called_once_with(raw_channel, interceptor) + def test_generate_channel_skips_intercept_channel_when_no_interceptors(self): raw_channel = Mock(spec=Channel) config = Config(tls=False, cache=CacheType.DISABLED)