diff --git a/core/pkg/service/ofrep/models.go b/core/pkg/service/ofrep/models.go index e7ff02849..2f656ae2e 100644 --- a/core/pkg/service/ofrep/models.go +++ b/core/pkg/service/ofrep/models.go @@ -22,6 +22,24 @@ type EvaluationSuccess struct { type BulkEvaluationResponse struct { Flags []interface{} `json:"flags"` Metadata model.Metadata `json:"metadata"` + // EventStreams advertises SSE endpoints clients can subscribe to for change + // notifications, per OpenFeature protocol ADR-0008. Omitted when SSE is disabled. + EventStreams []EventStream `json:"eventStreams,omitempty"` +} + +// EventStream describes a Server-Sent Events endpoint a client can subscribe to in order to +// be notified (via a `refetchEvaluation` event) when the flag configuration changes. +type EventStream struct { + Type string `json:"type"` + Endpoint *EventStreamEndpoint `json:"endpoint"` + InactivityDelaySec int `json:"inactivityDelaySec,omitempty"` +} + +// EventStreamEndpoint is the ADR-0008 structured form of an event-stream location. Origin is +// optional; when omitted the client resolves RequestUri against its OFREP base URL origin. +type EventStreamEndpoint struct { + Origin string `json:"origin,omitempty"` + RequestUri string `json:"requestUri"` } type EvaluationError struct { @@ -54,8 +72,8 @@ func BulkEvaluationResponseFrom(resolutions []evaluator.AnyValue, metadata model } return BulkEvaluationResponse{ - evaluations, - metadata, + Flags: evaluations, + Metadata: metadata, } } diff --git a/core/pkg/store/query.go b/core/pkg/store/query.go index 22dcaa558..fdd133bd4 100644 --- a/core/pkg/store/query.go +++ b/core/pkg/store/query.go @@ -83,6 +83,12 @@ func (s Selector) WithSource(source string) Selector { return s.withIndex(source func (s Selector) WithFlagSetId(id string) Selector { return s.withIndex(flagSetIdIndex, id) } func (s Selector) withKey(key string) Selector { return s.withIndex(keyIndex, key) } +// FlagSetId returns the flagSetId constraint of the selector, or "" if none is set. +func (s Selector) FlagSetId() string { return s.indexMap[flagSetIdIndex] } + +// Source returns the source constraint of the selector, or "" if none is set. +func (s Selector) Source() string { return s.indexMap[sourceIndex] } + func (s Selector) withIndex(key, value string) Selector { m := maps.Clone(s.indexMap) if m == nil { diff --git a/docs/reference/flagd-cli/flagd_start.md b/docs/reference/flagd-cli/flagd_start.md index 392cd24d9..4d98e9b4a 100644 --- a/docs/reference/flagd-cli/flagd_start.md +++ b/docs/reference/flagd-cli/flagd_start.md @@ -24,6 +24,9 @@ flagd start [flags] -R, --max-request-header int Maximum allowed request header size in bytes. Requests exceeding this are rejected with HTTP 431. Set to 0 to use Go's built-in default (1 MiB). WARNING: setting a very large or zero value may allow memory exhaustion from oversized headers. (default 1000000) -t, --metrics-exporter string Set the metrics exporter. Default(if unset) is Prometheus. Can be override to otel - OpenTelemetry metric exporter. Overriding to otel require otelCollectorURI to be present -r, --ofrep-port int32 ofrep service port (default 8016) + --ofrep-sse-enabled Enable the OFREP SSE change-notification endpoint (ADR-0008) at /ofrep/v1/sse/{channel} on the ofrep port, where the channel is a selector expression. Defaults to true. (default true) + --ofrep-sse-inactivity-delay int Inactivity delay (seconds) advertised to OFREP SSE clients in the eventStreams block. Clients close idle connections after this. Defaults to 120. (default 120) + --ofrep-sse-public-url string Origin (scheme://host) advertised as the OFREP SSE eventStreams endpoint.origin. Omitted when empty, so clients resolve the requestUri against the OFREP base URL. Set when flagd is behind a proxy. -A, --otel-ca-path string tls certificate authority path to use with OpenTelemetry collector -D, --otel-cert-path string tls certificate path to use with OpenTelemetry collector -o, --otel-collector-uri string Set the grpc URI of the OpenTelemetry collector for flagd runtime. If unset, the collector setup will be ignored and traces will not be exported. diff --git a/flagd/cmd/start.go b/flagd/cmd/start.go index e498304e1..d60faa832 100644 --- a/flagd/cmd/start.go +++ b/flagd/cmd/start.go @@ -23,6 +23,9 @@ const ( managementPortFlagName = "management-port" metricsExporter = "metrics-exporter" ofrepPortFlagName = "ofrep-port" + ofrepSSEEnabledFlagName = "ofrep-sse-enabled" + ofrepSSEInactivityFlagName = "ofrep-sse-inactivity-delay" + ofrepSSEPublicURLFlagName = "ofrep-sse-public-url" otelCollectorURI = "otel-collector-uri" otelCertPathFlagName = "otel-cert-path" otelKeyPathFlagName = "otel-key-path" @@ -58,6 +61,9 @@ func init() { flags.Int32P(syncPortFlagName, "g", 8015, "gRPC Sync port") flags.Int32P(ofrepPortFlagName, "r", 8016, "ofrep service port") + flags.Bool(ofrepSSEEnabledFlagName, true, "Enable the OFREP SSE change-notification endpoint (ADR-0008) at /ofrep/v1/sse/{channel} on the ofrep port, where the channel is a selector expression. Defaults to true.") + flags.Int(ofrepSSEInactivityFlagName, 120, "Inactivity delay (seconds) advertised to OFREP SSE clients in the eventStreams block. Clients close idle connections after this. Defaults to 120.") + flags.String(ofrepSSEPublicURLFlagName, "", "Origin (scheme://host) advertised as the OFREP SSE eventStreams endpoint.origin. Omitted when empty, so clients resolve the requestUri against the OFREP base URL. Set when flagd is behind a proxy.") flags.StringP(socketPathFlagName, "d", "", "Flagd unix socket path. "+ "With grpc the evaluations service will become available on this address. "+ "With http(s) the grpc-gateway proxy will use this address internally.") @@ -123,6 +129,9 @@ func bindFlags(flags *pflag.FlagSet) { _ = viper.BindPFlag(syncPortFlagName, flags.Lookup(syncPortFlagName)) _ = viper.BindPFlag(syncSocketPathFlagName, flags.Lookup(syncSocketPathFlagName)) _ = viper.BindPFlag(ofrepPortFlagName, flags.Lookup(ofrepPortFlagName)) + _ = viper.BindPFlag(ofrepSSEEnabledFlagName, flags.Lookup(ofrepSSEEnabledFlagName)) + _ = viper.BindPFlag(ofrepSSEInactivityFlagName, flags.Lookup(ofrepSSEInactivityFlagName)) + _ = viper.BindPFlag(ofrepSSEPublicURLFlagName, flags.Lookup(ofrepSSEPublicURLFlagName)) _ = viper.BindPFlag(contextValueFlagName, flags.Lookup(contextValueFlagName)) _ = viper.BindPFlag(headerToContextKeyFlagName, flags.Lookup(headerToContextKeyFlagName)) _ = viper.BindPFlag(streamDeadlineFlagName, flags.Lookup(streamDeadlineFlagName)) @@ -201,6 +210,9 @@ var startCmd = &cobra.Command{ MetricExporter: viper.GetString(metricsExporter), ManagementPort: viper.GetUint16(managementPortFlagName), OfrepServicePort: viper.GetUint16(ofrepPortFlagName), + OfrepSSEEnabled: viper.GetBool(ofrepSSEEnabledFlagName), + OfrepSSEInactivityDel: viper.GetInt(ofrepSSEInactivityFlagName), + OfrepSSEPublicURL: viper.GetString(ofrepSSEPublicURLFlagName), OtelCollectorURI: viper.GetString(otelCollectorURI), OtelCertPath: viper.GetString(otelCertPathFlagName), OtelKeyPath: viper.GetString(otelKeyPathFlagName), diff --git a/flagd/go.mod b/flagd/go.mod index 3ce5f74af..6f40cebfd 100644 --- a/flagd/go.mod +++ b/flagd/go.mod @@ -9,6 +9,7 @@ require ( connectrpc.com/connect v1.19.1 github.com/dimiro1/banner v1.1.0 github.com/gorilla/mux v1.8.1 + github.com/launchdarkly/eventsource v1.11.0 github.com/mattn/go-colorable v0.1.14 github.com/open-feature/flagd/core v0.15.6 github.com/prometheus/client_golang v1.23.2 diff --git a/flagd/go.sum b/flagd/go.sum index b1a4529cd..fd196ad2a 100644 --- a/flagd/go.sum +++ b/flagd/go.sum @@ -255,6 +255,10 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc= github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw= +github.com/launchdarkly/eventsource v1.11.0 h1:aAdvh2XmtXA17QsRFL0XKHURMqhxg7J+CceQmhSzBas= +github.com/launchdarkly/eventsource v1.11.0/go.mod h1:dU+rZxkPOlGPsyJPpiDqiepAcFwIITDUClY9+A6RrMw= +github.com/launchdarkly/go-test-helpers/v3 v3.1.0 h1:E3bxJMzMoA+cJSF3xxtk2/chr1zshl1ZWa0/oR+8bvg= +github.com/launchdarkly/go-test-helpers/v3 v3.1.0/go.mod h1:Ake5+hZFS/DmIGKx/cizhn5W9pGA7pplcR7xCxWiLIo= github.com/mailru/easyjson v0.7.7 h1:UGYAvKxe3sBsEDzO8ZeWOSlIQfWFlxbzLZe7hwFURr0= github.com/mailru/easyjson v0.7.7/go.mod h1:xzfreul335JAWq5oZzymOObrkdz5UnU4kGfJJLY9Nlc= github.com/mattn/go-colorable v0.1.4/go.mod h1:U0ppj6V5qS13XJ6of8GYAs25YV2eR4EVcfRqFIhoBtE= diff --git a/flagd/pkg/runtime/from_config.go b/flagd/pkg/runtime/from_config.go index 777a1af98..7b93e1011 100644 --- a/flagd/pkg/runtime/from_config.go +++ b/flagd/pkg/runtime/from_config.go @@ -27,6 +27,9 @@ type Config struct { MetricExporter string ManagementPort uint16 OfrepServicePort uint16 + OfrepSSEEnabled bool + OfrepSSEInactivityDel int + OfrepSSEPublicURL string OtelCollectorURI string OtelCertPath string OtelKeyPath string @@ -112,13 +115,16 @@ func FromConfig(logger *logger.Logger, version string, config Config) (*Runtime, recorder) // ofrep service - ofrepService, err := ofrep.NewOfrepService(jsonEvaluator, config.CORS, ofrep.SvcConfiguration{ + ofrepService, err := ofrep.NewOfrepService(jsonEvaluator, store, config.CORS, ofrep.SvcConfiguration{ Logger: logger.WithFields(zap.String("component", "OFREPService")), Port: config.OfrepServicePort, ServiceName: svcName, MetricsRecorder: recorder, MaxRequestBodyBytes: config.MaxRequestBodyBytes, MaxRequestHeaderBytes: config.MaxRequestHeaderBytes, + SSEEnabled: config.OfrepSSEEnabled, + SSEInactivityDelaySec: config.OfrepSSEInactivityDel, + SSEPublicURL: config.OfrepSSEPublicURL, }, config.ContextValues, config.HeaderToContextKeyMappings, diff --git a/flagd/pkg/service/flag-evaluation/ofrep/handler.go b/flagd/pkg/service/flag-evaluation/ofrep/handler.go index bd4b8a344..f90708828 100644 --- a/flagd/pkg/service/flag-evaluation/ofrep/handler.go +++ b/flagd/pkg/service/flag-evaluation/ofrep/handler.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "net/http" + "strings" "github.com/gorilla/mux" "github.com/open-feature/flagd/core/pkg/evaluator" @@ -16,6 +17,7 @@ import ( "github.com/open-feature/flagd/core/pkg/telemetry" "github.com/open-feature/flagd/flagd/pkg/service" evalservice "github.com/open-feature/flagd/flagd/pkg/service/flag-evaluation" + "github.com/open-feature/flagd/flagd/pkg/service/flag-evaluation/ofrep/sse" metricsmw "github.com/open-feature/flagd/flagd/pkg/service/middleware/metrics" "github.com/rs/xid" "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" @@ -29,6 +31,14 @@ const ( bulkEvaluation = "/ofrep/v1/evaluate/{path:flags\\/|flags}" ) +// configVersioner resolves the current config ETag / last-modified time for an SSE channel so +// the bulk handler can serve conditional (ETag/304) responses consistent with the SSE stream. +// Implemented by the OFREP SSE change tracker; nil when SSE is disabled. ok is false whenever no +// stream for the channel is live, which is the normal state before a client connects. +type configVersioner interface { + Version(channel string) (etag string, lastModified int64, ok bool) +} + type handler struct { Logger *logger.Logger evaluator evaluator.IEvaluator @@ -36,6 +46,20 @@ type handler struct { headerToContextKeyMappings map[string]string metricsRecorder telemetry.IMetricsRecorder tracer trace.Tracer + + versioner configVersioner + sseEnabled bool + sseInactivityDelaySec int + ssePublicURL string +} + +// SSEConfig carries the SSE advertisement settings the bulk handler needs to expose the +// `eventStreams` block and conditional-evaluation ETags. +type SSEConfig struct { + Enabled bool + Versioner configVersioner + InactivityDelaySec int + PublicURL string } func NewOfrepHandler( @@ -45,6 +69,7 @@ func NewOfrepHandler( headerToContextKeyMappings map[string]string, metricsRecorder telemetry.IMetricsRecorder, serviceName string, + sseCfg SSEConfig, ) http.Handler { h := handler{ Logger: logger, @@ -53,6 +78,10 @@ func NewOfrepHandler( headerToContextKeyMappings: headerToContextKeyMappings, metricsRecorder: metricsRecorder, tracer: otel.Tracer("flagd.ofrep.v1"), + versioner: sseCfg.Versioner, + sseEnabled: sseCfg.Enabled, + sseInactivityDelaySec: sseCfg.InactivityDelaySec, + ssePublicURL: sseCfg.PublicURL, } router := mux.NewRouter() @@ -143,6 +172,14 @@ func (h *handler) HandleBulkEvaluation(w http.ResponseWriter, r *http.Request) { } ctx := context.WithValue(r.Context(), store.SelectorContextKey{}, selector) + // Conditional evaluation (ADR-0008): short-circuit with 304 when the client already holds + // the current config version. + lastModified, notModified := h.applyConditionalETag(w, r, selectorExpression) + if notModified { + w.WriteHeader(http.StatusNotModified) + return + } + evaluations, metadata, err := h.evaluator.ResolveAllValues(ctx, requestID, evaluationContext) if h.metricsRecorder != nil { for _, evaluation := range evaluations { @@ -155,9 +192,84 @@ func (h *handler) HandleBulkEvaluation(w http.ResponseWriter, r *http.Request) { res := ofrep.BulkEvaluationContextErrorFrom(model.GeneralErrorCode, fmt.Sprintf("Bulk evaluation failed. Tracking ID: %s", requestID)) h.writeJSONToResponse(http.StatusInternalServerError, res, w) - } else { - h.writeJSONToResponse(http.StatusOK, ofrep.BulkEvaluationResponseFrom(evaluations, metadata), w) + return + } + + h.writeJSONToResponse(http.StatusOK, + h.bulkResponse(selectorExpression, evaluations, metadata, lastModified), w) +} + +// applyConditionalETag resolves the current config version for the selector, sets the ETag +// response header, and reports the lastModified time plus whether the request can be answered +// with 304 Not Modified. It is a no-op (returns notModified=false) when SSE/versioning is off. +func (h *handler) applyConditionalETag(w http.ResponseWriter, r *http.Request, channel string) (lastModified int64, notModified bool) { + if h.versioner == nil { + return 0, false } + etag, lastModified, ok := h.versioner.Version(channel) + if !ok || etag == "" { + return lastModified, false + } + w.Header().Set("ETag", quoteETag(etag)) + + if trigger := r.URL.Query().Get(flagConfigEtagParam); trigger != "" { + h.Logger.Debug(fmt.Sprintf("bulk refetch triggered by %s=%s", flagConfigEtagParam, trigger)) + } + + clientCacheETag := r.Header.Get("If-None-Match") + return lastModified, clientCacheETag != "" && normalizeETag(clientCacheETag) == etag +} + +// bulkResponse assembles the OFREP bulk response, adding the ADR-0008 eventStreams block and +// lastModified metadata when SSE is enabled. +func (h *handler) bulkResponse(selectorExpression string, evaluations []evaluator.AnyValue, metadata model.Metadata, lastModified int64) ofrep.BulkEvaluationResponse { + response := ofrep.BulkEvaluationResponseFrom(evaluations, metadata) + if !h.sseEnabled { + return response + } + response.EventStreams = h.eventStreams(selectorExpression) + if lastModified > 0 { + if response.Metadata == nil { + response.Metadata = model.Metadata{} + } + response.Metadata["flagConfigLastModified"] = lastModified + } + return response +} + +// eventStreams builds the ADR-0008 eventStreams advertisement pointing OFREP clients back at +// this flagd's SSE endpoint. The advertised channel is the request's own selector expression, +// carried as the final path segment, so the stream covers exactly the flags the client just +// evaluated. +// +// It uses the structured `endpoint` form and omits origin unless a public URL is configured, so +// the client resolves the requestUri against the OFREP base URL it is already talking to. +func (h *handler) eventStreams(selectorExpression string) []ofrep.EventStream { + requestURI := sse.ChannelPath(ssePath, selectorExpression) + + return []ofrep.EventStream{{ + Type: "sse", + InactivityDelaySec: h.sseInactivityDelaySec, + Endpoint: &ofrep.EventStreamEndpoint{ + Origin: strings.TrimSuffix(h.ssePublicURL, "/"), + RequestUri: requestURI, + }, + }} +} + +// flagConfigEtagParam is the ADR-0008 query parameter carrying the config version that an SSE +// refetch event advertised. It is trigger metadata only; the conditional response is decided by +// the If-None-Match header. See applyConditionalETag. +const flagConfigEtagParam = "flagConfigEtag" + +// normalizeETag strips optional surrounding quotes so quoted and unquoted forms compare equal. +func normalizeETag(etag string) string { + return strings.Trim(etag, `"`) +} + +// quoteETag wraps a bare ETag value in the double quotes required by the HTTP ETag header. +func quoteETag(etag string) string { + return `"` + etag + `"` } func (h *handler) writeJSONToResponse(status int, payload interface{}, w http.ResponseWriter) { diff --git a/flagd/pkg/service/flag-evaluation/ofrep/handler_test.go b/flagd/pkg/service/flag-evaluation/ofrep/handler_test.go index fdb01784b..a530b03cb 100644 --- a/flagd/pkg/service/flag-evaluation/ofrep/handler_test.go +++ b/flagd/pkg/service/flag-evaluation/ofrep/handler_test.go @@ -472,7 +472,7 @@ func TestHandlerRecordsSingleEvaluationMetrics(t *testing.T) { eval.EXPECT(). ResolveAsAnyValue(gomock.Any(), gomock.Any(), flagKey, gomock.Any()). Return(test.evaluation) - handler := NewOfrepHandler(logger.NewLogger(nil, false), eval, nil, nil, metrics, "flagd") + handler := NewOfrepHandler(logger.NewLogger(nil, false), eval, nil, nil, metrics, "flagd", SSEConfig{}) request := httptest.NewRequest(http.MethodPost, "/ofrep/v1/evaluate/flags/"+flagKey, nil) response := httptest.NewRecorder() @@ -490,7 +490,7 @@ func TestHandlerRecordsEachBulkEvaluationMetric(t *testing.T) { eval := mock.NewMockIEvaluator(gomock.NewController(t)) eval.EXPECT().ResolveAllValues(gomock.Any(), gomock.Any(), gomock.Any()). Return(evaluations, model.Metadata{}, nil) - handler := NewOfrepHandler(logger.NewLogger(nil, false), eval, nil, nil, metrics, "flagd") + handler := NewOfrepHandler(logger.NewLogger(nil, false), eval, nil, nil, metrics, "flagd", SSEConfig{}) request := httptest.NewRequest(http.MethodPost, "/ofrep/v1/evaluate/flags", nil) response := httptest.NewRecorder() diff --git a/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service.go b/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service.go index 9a8994469..2e3a2febb 100644 --- a/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service.go +++ b/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service.go @@ -9,11 +9,16 @@ import ( "github.com/open-feature/flagd/core/pkg/evaluator" "github.com/open-feature/flagd/core/pkg/logger" + "github.com/open-feature/flagd/core/pkg/store" "github.com/open-feature/flagd/core/pkg/telemetry" + "github.com/open-feature/flagd/flagd/pkg/service/flag-evaluation/ofrep/sse" corsmw "github.com/open-feature/flagd/flagd/pkg/service/middleware/cors" "golang.org/x/sync/errgroup" ) +// ssePath is the endpoint OFREP clients subscribe to for change notifications (ADR-0008). +const ssePath = "/ofrep/v1/sse" + type IOfrepService interface { // Start the OFREP service with context for shutdown Start(context.Context) error @@ -26,30 +31,68 @@ type SvcConfiguration struct { MetricsRecorder telemetry.IMetricsRecorder MaxRequestBodyBytes int64 MaxRequestHeaderBytes int64 + + // SSE (ADR-0008) settings + SSEEnabled bool + SSEInactivityDelaySec int + SSEPublicURL string } type Service struct { logger *logger.Logger port uint16 server *http.Server + sse *sse.Service } func NewOfrepService( - evaluator evaluator.IEvaluator, origins []string, cfg SvcConfiguration, contextValues map[string]any, headerToContextKeyMappings map[string]string, + evaluator evaluator.IEvaluator, flagStore store.IStore, origins []string, cfg SvcConfiguration, contextValues map[string]any, headerToContextKeyMappings map[string]string, ) (*Service, error) { corsMiddleware := corsmw.New(origins) - var h http.Handler = NewOfrepHandler( + // SSE requires a flag store to watch; without one we cannot serve or advertise it, so treat + // a missing store as SSE disabled rather than crashing later in the tracker goroutine. + sseEnabled := cfg.SSEEnabled + if sseEnabled && flagStore == nil { + cfg.Logger.Warn("OFREP SSE requested but no flag store was provided; disabling SSE") + sseEnabled = false + } + + var sseService *sse.Service + sseCfg := SSEConfig{ + Enabled: sseEnabled, + InactivityDelaySec: cfg.SSEInactivityDelaySec, + PublicURL: cfg.SSEPublicURL, + } + if sseEnabled { + sseService = sse.New(flagStore, sse.Config{Logger: cfg.Logger}) + sseCfg.Versioner = sseService.Tracker() + } + + ofrepHandler := NewOfrepHandler( cfg.Logger, evaluator, contextValues, headerToContextKeyMappings, cfg.MetricsRecorder, cfg.ServiceName, + sseCfg, ) + + // Route the long-lived SSE stream separately from the request/response evaluate routes so + // the request-body limit is not applied to the stream. Everything else falls through to the + // existing OFREP handler. + mux := http.NewServeMux() + var evaluateHandler http.Handler = ofrepHandler if cfg.MaxRequestBodyBytes > 0 { - h = http.MaxBytesHandler(h, cfg.MaxRequestBodyBytes) + evaluateHandler = http.MaxBytesHandler(evaluateHandler, cfg.MaxRequestBodyBytes) + } + if sseService != nil { + sseService.Register(mux, ssePath) } + mux.Handle("/", evaluateHandler) + + var h http.Handler = mux h = corsMiddleware.Handler(h) server := http.Server{ @@ -65,6 +108,7 @@ func NewOfrepService( logger: cfg.Logger, port: cfg.Port, server: &server, + sse: sseService, }, nil } @@ -81,6 +125,12 @@ func (s Service) Start(ctx context.Context) error { return nil }) + if s.sse != nil { + group.Go(func() error { + return s.sse.Start(gCtx) + }) + } + group.Go(func() error { <-gCtx.Done() s.logger.Info("shutting down ofrep service") diff --git a/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service_test.go b/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service_test.go index 90682a4c4..51e2ed462 100644 --- a/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service_test.go +++ b/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service_test.go @@ -37,7 +37,7 @@ func TestOfrepServiceStartStop(t *testing.T) { MetricsRecorder: &telemetry.NoopMetricsRecorder{}, } - service, err := NewOfrepService(eval, []string{"*"}, cfg, nil, nil) + service, err := NewOfrepService(eval, nil, []string{"*"}, cfg, nil, nil) if err != nil { t.Fatalf(errCreateOfrepService, err) } @@ -193,7 +193,7 @@ func startOfrepService(t *testing.T, cfg SvcConfiguration) (*Service, uint16) { t.Helper() eval := mock.NewMockIEvaluator(gomock.NewController(t)) - service, err := NewOfrepService(eval, []string{"*"}, cfg, nil, nil) + service, err := NewOfrepService(eval, nil, []string{"*"}, cfg, nil, nil) if err != nil { t.Fatalf(errCreateOfrepService, err) } diff --git a/flagd/pkg/service/flag-evaluation/ofrep/sse/event.go b/flagd/pkg/service/flag-evaluation/ofrep/sse/event.go new file mode 100644 index 000000000..811402d74 --- /dev/null +++ b/flagd/pkg/service/flag-evaluation/ofrep/sse/event.go @@ -0,0 +1,36 @@ +package sse + +import "encoding/json" + +// refetchEventType is the only event data type defined by OpenFeature protocol ADR-0008. +const refetchEventType = "refetchEvaluation" + +// eventName is the SSE `event:` field, using the ADR-0008 "message" envelope. +const eventName = "message" + +type refetchPayload struct { + Type string `json:"type"` + Etag string `json:"etag,omitempty"` + LastModified int64 `json:"lastModified,omitempty"` +} + +// refetchEvent implements eventsource.Event, telling subscribed OFREP clients that the flag +// configuration for their channel changed and they should re-run bulk evaluation. +type refetchEvent struct { + id string + data string +} + +func newRefetchEvent(id, etag string, lastModified int64) refetchEvent { + payload := refetchPayload{ + Type: refetchEventType, + Etag: etag, + LastModified: lastModified, + } + data, _ := json.Marshal(payload) + return refetchEvent{id: id, data: string(data)} +} + +func (e refetchEvent) Id() string { return e.id } +func (e refetchEvent) Event() string { return eventName } +func (e refetchEvent) Data() string { return e.data } diff --git a/flagd/pkg/service/flag-evaluation/ofrep/sse/handler.go b/flagd/pkg/service/flag-evaluation/ofrep/sse/handler.go new file mode 100644 index 000000000..0471c6808 --- /dev/null +++ b/flagd/pkg/service/flag-evaluation/ofrep/sse/handler.go @@ -0,0 +1,39 @@ +package sse + +import ( + "net/http" + + "github.com/open-feature/flagd/core/pkg/store" +) + +// channelPathVar is the path wildcard carrying the selector expression the client wants change +// notifications for, using the same syntax as the Flagd-Selector header. An empty channel (the +// bare stream path) selects every flag. See Service.Register for the routes it is bound to. +const channelPathVar = "channel" + +// Handler resolves the request's selector, takes a reference on the matching subscription so the +// store watch stays alive, and streams events until the client disconnects. +// +// It must be mounted through Service.Register: the channel is read from the request path, so the +// route has to declare the channel wildcard. +func (svc *Service) Handler() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // already percent-decoded by the mux, so a source containing "/" arrives intact + channel := r.PathValue(channelPathVar) + selector, err := store.NewSelector(channel) + if err != nil { + // not echoing the expression back: it is unescaped client input + http.Error(w, "invalid selector in the channel path segment", http.StatusBadRequest) + return + } + + release, err := svc.tracker.Subscribe(channel, selector) + if err != nil { + http.Error(w, "unable to subscribe to flag changes", http.StatusServiceUnavailable) + return + } + defer release() + + svc.es.Handler(channel).ServeHTTP(w, r) + }) +} diff --git a/flagd/pkg/service/flag-evaluation/ofrep/sse/service.go b/flagd/pkg/service/flag-evaluation/ofrep/sse/service.go new file mode 100644 index 000000000..896fd890e --- /dev/null +++ b/flagd/pkg/service/flag-evaluation/ofrep/sse/service.go @@ -0,0 +1,97 @@ +package sse + +import ( + "context" + "net/http" + "net/url" + "time" + + "github.com/launchdarkly/eventsource" + "github.com/open-feature/flagd/core/pkg/logger" + "github.com/open-feature/flagd/core/pkg/store" +) + +// defaultHeartbeatInterval must stay well below the advertised inactivity delay so idle +// connections are not closed by clients or proxies. +const defaultHeartbeatInterval = 30 * time.Second + +type Config struct { + Logger *logger.Logger + HeartbeatInterval time.Duration +} + +// Service serves the OFREP SSE stream. It owns an eventsource server, the change Tracker that +// drives it, and a heartbeat loop that keeps idle connections alive. +type Service struct { + logger *logger.Logger + es *eventsource.Server + tracker *Tracker + heartbeatInterval time.Duration +} + +// New builds an SSE Service backed by the shared flag store. +func New(s store.IStore, cfg Config) *Service { + es := eventsource.NewServer() + // CORS is handled by the OFREP server's middleware; avoid emitting duplicate headers. + es.AllowCORS = false + es.ReplayAll = false + + heartbeat := cfg.HeartbeatInterval + if heartbeat <= 0 { + heartbeat = defaultHeartbeatInterval + } + + return &Service{ + logger: cfg.Logger, + es: es, + tracker: NewTracker(cfg.Logger, s, es), + heartbeatInterval: heartbeat, + } +} + +// Tracker exposes the change tracker so the OFREP bulk handler can resolve config versions +// for conditional (ETag/304) evaluation. +func (svc *Service) Tracker() *Tracker { return svc.tracker } + +// Register mounts the stream on mux as {prefix}/{channel}, where the channel is a selector +// expression. The bare prefix is registered too, so subscribing to every flag does not depend on +// a trailing-slash redirect. +func (svc *Service) Register(mux *http.ServeMux, prefix string) { + h := svc.Handler() + mux.Handle(prefix, h) + mux.Handle(prefix+"/{"+channelPathVar+"}", h) +} + +// ChannelPath returns the stream path for a channel under prefix. The selector expression is a +// single path segment, so it is escaped: source selectors routinely contain "/". +func ChannelPath(prefix, channel string) string { + if channel == "" { + return prefix + } + return prefix + "/" + url.PathEscape(channel) +} + +// Start runs the heartbeat loop until ctx is cancelled, then shuts down. It blocks, so it is +// intended to run in its own goroutine. +func (svc *Service) Start(ctx context.Context) error { + ticker := time.NewTicker(svc.heartbeatInterval) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + // Order matters: the tracker's watch goroutines are the only publishers, and + // eventsource.Server.Publish blocks forever once the server is closed. + svc.tracker.Close() + svc.es.Close() + if svc.logger != nil { + svc.logger.Info("shutting down ofrep sse service") + } + return nil + case <-ticker.C: + if channels := svc.tracker.Channels(); len(channels) > 0 { + svc.es.PublishComment(channels, "keep-alive") + } + } + } +} diff --git a/flagd/pkg/service/flag-evaluation/ofrep/sse/service_test.go b/flagd/pkg/service/flag-evaluation/ofrep/sse/service_test.go new file mode 100644 index 000000000..b936891a1 --- /dev/null +++ b/flagd/pkg/service/flag-evaluation/ofrep/sse/service_test.go @@ -0,0 +1,153 @@ +package sse + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/launchdarkly/eventsource" + "github.com/open-feature/flagd/core/pkg/logger" + "github.com/open-feature/flagd/core/pkg/model" + "github.com/open-feature/flagd/core/pkg/store" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// testSSEPath mirrors the prefix the OFREP service mounts the stream on. +const testSSEPath = "/ofrep/v1/sse" + +// newStreamServer mounts the service the way the OFREP service does, so requests exercise the +// real routing: the channel is the final path segment. +func newStreamServer(t *testing.T, svc *Service) *httptest.Server { + t.Helper() + mux := http.NewServeMux() + svc.Register(mux, testSSEPath) + ts := httptest.NewServer(mux) + t.Cleanup(ts.Close) + return ts +} + +// streamURL builds the stream URL for a channel on ts. +func streamURL(ts *httptest.Server, channel string) string { + return ts.URL + ChannelPath(testSSEPath, channel) +} + +// TestService_PublishesRefetchOnChange verifies the end-to-end wiring: a client subscribed with +// a flagSetId selector receives an ADR-0008 refetchEvaluation event when that flag set changes. +func TestService_PublishesRefetchOnChange(t *testing.T) { + log := logger.NewLogger(nil, false) + s, err := store.NewStore(log, []string{"src1"}) + require.NoError(t, err) + + // long heartbeat so keep-alive comments do not interfere with the assertion + svc := New(s, Config{Logger: log, HeartbeatInterval: time.Hour}) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go func() { _ = svc.Start(ctx) }() + + ts := newStreamServer(t, svc) + + // the channel token is a selector expression, same syntax as Flagd-Selector + stream, err := eventsource.SubscribeWithURL(streamURL(ts, "flagSetId=fs1")) + require.NoError(t, err) + defer stream.Close() + + // eventsource writes the 200 before registering, so Subscribe returning is not enough + time.Sleep(100 * time.Millisecond) + + s.Update("src1", []model.Flag{testFlag("fs1", "a", "on")}, model.Metadata{"flagSetId": "fs1"}, false) + + select { + case ev := <-stream.Events: + assert.Equal(t, eventName, ev.Event()) + assert.Contains(t, ev.Data(), refetchEventType) + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for refetch event on channel fs1") + } +} + +// TestService_SelectorScopedNotifications asserts a client is woken only for changes to the +// flags its selector selects, including a source= selector, which previously fell back to the +// catch-all channel and was woken by every change. +func TestService_SelectorScopedNotifications(t *testing.T) { + log := logger.NewLogger(nil, false) + s, err := store.NewStore(log, []string{"src1", "src2"}) + require.NoError(t, err) + + s.Update("src1", []model.Flag{testFlag("fs1", "a", "on")}, model.Metadata{"flagSetId": "fs1"}, false) + s.Update("src2", []model.Flag{testFlag("fs2", "b", "on")}, model.Metadata{"flagSetId": "fs2"}, false) + + svc := New(s, Config{Logger: log, HeartbeatInterval: time.Hour}) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go func() { _ = svc.Start(ctx) }() + + ts := newStreamServer(t, svc) + + bySource, err := eventsource.SubscribeWithURL(streamURL(ts, "source=src1")) + require.NoError(t, err) + defer bySource.Close() + + byFlagSet, err := eventsource.SubscribeWithURL(streamURL(ts, "flagSetId=fs2")) + require.NoError(t, err) + defer byFlagSet.Close() + + time.Sleep(100 * time.Millisecond) + + s.Update("src2", []model.Flag{testFlag("fs2", "b", "off")}, model.Metadata{"flagSetId": "fs2"}, false) + + select { + case ev := <-byFlagSet.Events: + assert.Contains(t, ev.Data(), refetchEventType) + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for a refetch event on the flagSetId=fs2 stream") + } + + select { + case ev := <-bySource.Events: + t.Fatalf("source=src1 must not be notified about a src2 change, got %q", ev.Data()) + case <-time.After(500 * time.Millisecond): + } +} + +func TestService_Handler_InvalidSelectorReturns400(t *testing.T) { + log := logger.NewLogger(nil, false) + s, err := store.NewStore(log, []string{"src1"}) + require.NoError(t, err) + + svc := New(s, Config{Logger: log, HeartbeatInterval: time.Hour}) + ts := newStreamServer(t, svc) + + resp, err := http.Get(streamURL(ts, "bogusKey=1")) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusBadRequest, resp.StatusCode) + assert.Empty(t, svc.Tracker().Channels(), "a rejected request must not create a subscription") +} + +// TestService_ChannelComesFromPath pins the routing: the channel is the final path segment, the +// bare path is the catch-all, and a selector carrying a "/" survives the round trip escaped. +func TestService_ChannelComesFromPath(t *testing.T) { + log := logger.NewLogger(nil, false) + s, err := store.NewStore(log, []string{"./mySource"}) + require.NoError(t, err) + + svc := New(s, Config{Logger: log, HeartbeatInterval: time.Hour}) + ts := newStreamServer(t, svc) + + for _, channel := range []string{"", "flagSetId=fs1", "source=./mySource"} { + stream, err := eventsource.SubscribeWithURL(streamURL(ts, channel)) + require.NoError(t, err, "channel %q", channel) + defer stream.Close() + } + + assert.Eventually(t, func() bool { + return len(svc.Tracker().Channels()) == 3 + }, 3*time.Second, 5*time.Millisecond) + assert.ElementsMatch(t, []string{"", "flagSetId=fs1", "source=./mySource"}, svc.Tracker().Channels(), + "the subscription must key off the decoded path segment, not its escaped form") +} diff --git a/flagd/pkg/service/flag-evaluation/ofrep/sse/tracker.go b/flagd/pkg/service/flag-evaluation/ofrep/sse/tracker.go new file mode 100644 index 000000000..5b9cfa8da --- /dev/null +++ b/flagd/pkg/service/flag-evaluation/ofrep/sse/tracker.go @@ -0,0 +1,229 @@ +package sse + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "sort" + "strconv" + "sync" + "sync/atomic" + "time" + + "github.com/launchdarkly/eventsource" + "github.com/open-feature/flagd/core/pkg/logger" + "github.com/open-feature/flagd/core/pkg/model" + "github.com/open-feature/flagd/core/pkg/store" + "go.uber.org/zap" +) + +type version struct { + etag string + lastModified int64 +} + +// subscription is a single store.Watch, shared by every client on the same channel. It is +// reachable through Tracker.subs only while refs > 0. +type subscription struct { + selector store.Selector // immutable after construction: the store retains &selector + cancel context.CancelFunc + + // ver is nil until the first snapshot lands, and again if the watch dies. + ver atomic.Pointer[version] + + refs int // guarded by Tracker.mu +} + +type Tracker struct { + logger *logger.Logger + store store.IStore + es *eventsource.Server + + mu sync.Mutex + subs map[string]*subscription + closed bool + + wg sync.WaitGroup +} + +func NewTracker(log *logger.Logger, s store.IStore, es *eventsource.Server) *Tracker { + return &Tracker{ + logger: log, + store: s, + es: es, + subs: map[string]*subscription{}, + } +} + +// Subscribe registers interest in a channel, starting a store watch for its selector if this is +// the first subscriber. The returned release must be called exactly once on disconnect. +func (t *Tracker) Subscribe(channel string, selector store.Selector) (func(), error) { + t.mu.Lock() + defer t.mu.Unlock() + + if t.closed { + return nil, fmt.Errorf("ofrep sse tracker is shutting down") + } + + sub, existing := t.subs[channel] + if !existing { + // Subscriptions are shared, so they outlive any single request and are torn down by + // release or Close rather than by a request context. + ctx, cancel := context.WithCancel(context.Background()) + sub = &subscription{selector: selector, cancel: cancel} + t.subs[channel] = sub + t.wg.Add(1) + + watcher := make(chan store.FlagQueryResult, 1) + t.store.Watch(ctx, &sub.selector, watcher) + go t.watch(ctx, channel, sub, watcher) + } + sub.refs++ + + var once sync.Once + return func() { once.Do(func() { t.release(channel, sub) }) }, nil +} + +// release drops one reference, tearing the subscription down when the last one goes. The map +// delete happens under the same lock as the refcount check and before cancel, so a concurrent +// Subscribe cannot attach to a dying subscription. +func (t *Tracker) release(channel string, sub *subscription) { + t.mu.Lock() + sub.refs-- + last := sub.refs <= 0 + if last && t.subs[channel] == sub { + delete(t.subs, channel) + } + t.mu.Unlock() + + if last { + sub.cancel() + } +} + +func (t *Tracker) watch(ctx context.Context, channel string, sub *subscription, watcher <-chan store.FlagQueryResult) { + defer t.wg.Done() + + first := true + // This loop must run until the channel closes: store.Watch's send does not select on + // ctx.Done(), so abandoning the channel would leak the store's goroutine. + for res := range watcher { + fp := fingerprint(&sub.selector, res.Flags) + prev := sub.ver.Load() + if prev != nil && prev.etag == fp { + // A no-op wakeup: an identical re-sync, or a coarser radix watch channel firing + // for a change outside this selector. + continue + } + sub.ver.CompareAndSwap(prev, &version{etag: fp, lastModified: time.Now().Unix()}) + + if first { + // seed only; the connecting client re-fetches unconditionally anyway (ADR-0008) + first = false + continue + } + t.publish(channel, sub) + } + + if ctx.Err() != nil { + return // our own teardown + } + + // The store hit a selector/iterator error, so this channel will never fire again. Stop + // serving its ETag, which would otherwise answer 304 from a frozen fingerprint forever. + sub.ver.Store(nil) + if t.logger != nil { + t.logger.Error("ofrep sse watch ended unexpectedly; its cached ETag is invalidated", zap.String("channel", channel)) + } +} + +// publish must be called only after the new version is stored, so a client that refetches the +// instant it receives the event cannot read the previous ETag and be served a stale 304. +func (t *Tracker) publish(channel string, sub *subscription) { + v := sub.ver.Load() + if v == nil { + return + } + + id := strconv.FormatInt(time.Now().UnixMilli(), 10) + t.es.Publish([]string{channel}, newRefetchEvent(id, v.etag, v.lastModified)) + if t.logger != nil { + t.logger.Debug("published refetch event", zap.String("channel", channel), zap.String("etag", v.etag)) + } +} + +// Version returns the current config ETag and last-modified time (unix seconds) for the channel +// matching the given selector expression. ok is false when no stream for it is live, so the +// caller skips conditional handling. +func (t *Tracker) Version(channel string) (etag string, lastModified int64, ok bool) { + t.mu.Lock() + sub, exists := t.subs[channel] + t.mu.Unlock() + + if !exists { + return "", 0, false + } + v := sub.ver.Load() + if v == nil { + return "", 0, false + } + return v.etag, v.lastModified, true +} + +// Channels returns the channels with at least one live subscriber. +func (t *Tracker) Channels() []string { + t.mu.Lock() + defer t.mu.Unlock() + + channels := make([]string, 0, len(t.subs)) + for channel := range t.subs { + channels = append(channels, channel) + } + return channels +} + +// Close stops accepting subscriptions and waits for every watch goroutine to exit. It must +// complete before the eventsource server is closed, since a publish after that blocks forever. +func (t *Tracker) Close() { + t.mu.Lock() + t.closed = true + subs := make([]*subscription, 0, len(t.subs)) + for _, sub := range t.subs { + subs = append(subs, sub) + } + t.subs = map[string]*subscription{} + t.mu.Unlock() + + for _, sub := range subs { + sub.cancel() + } + t.wg.Wait() +} + +// fingerprint produces a deterministic hash of the result set a channel's watch resolves to. +func fingerprint(selector *store.Selector, flags []model.Flag) string { + digests := make([]string, 0, len(flags)) + for _, f := range flags { + h := sha256.New() + h.Write([]byte(f.Key)) + h.Write([]byte{0}) + h.Write([]byte(f.Source)) + h.Write([]byte{0}) + b, _ := json.Marshal(f) + h.Write(b) + digests = append(digests, string(h.Sum(nil))) + } + sort.Strings(digests) + + h := sha256.New() + // the selector path scopes the hash to this watch, so two channels resolving to the same + // flags never share an ETag unless they come from the same selector source. + h.Write([]byte(selector.ToLogString())) + h.Write([]byte{0}) + for _, d := range digests { + h.Write([]byte(d)) + } + return hex.EncodeToString(h.Sum(nil)) +} diff --git a/flagd/pkg/service/flag-evaluation/ofrep/sse/tracker_test.go b/flagd/pkg/service/flag-evaluation/ofrep/sse/tracker_test.go new file mode 100644 index 000000000..cf99d3ae0 --- /dev/null +++ b/flagd/pkg/service/flag-evaluation/ofrep/sse/tracker_test.go @@ -0,0 +1,264 @@ +package sse + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/launchdarkly/eventsource" + + "github.com/open-feature/flagd/core/pkg/logger" + "github.com/open-feature/flagd/core/pkg/model" + "github.com/open-feature/flagd/core/pkg/store" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func testFlag(flagSetID, key, defaultVariant string) model.Flag { + return model.Flag{ + Key: key, + FlagSetId: flagSetID, + State: "ENABLED", + DefaultVariant: defaultVariant, + Variants: map[string]any{"on": true, "off": false}, + Source: "src1", + } +} + +func mustSelector(t *testing.T, expr string) store.Selector { + t.Helper() + s, err := store.NewSelector(expr) + require.NoError(t, err) + return s +} + +func newTestStore(t *testing.T) *store.Store { + t.Helper() + s, err := store.NewStore(logger.NewLogger(nil, false), []string{"src1", "src2"}) + require.NoError(t, err) + return s +} + +// newTracker builds a Tracker over a real eventsource server, so publishes are exercised. +func newTracker(t *testing.T, s store.IStore) *Tracker { + t.Helper() + es := eventsource.NewServer() + t.Cleanup(es.Close) + + tr := NewTracker(logger.NewLogger(nil, false), s, es) + t.Cleanup(tr.Close) + return tr +} + +// subscribe registers a channel and ties its release to the test cleanup. +func subscribe(t *testing.T, tr *Tracker, expr string) { + t.Helper() + release, err := tr.Subscribe(expr, mustSelector(t, expr)) + require.NoError(t, err) + t.Cleanup(release) +} + +// versionOf polls until the channel has a seeded ETag, since seeding is asynchronous. +func versionOf(t *testing.T, tr *Tracker, channel string) (string, int64) { + t.Helper() + var etag string + var lastModified int64 + require.Eventually(t, func() bool { + var ok bool + etag, lastModified, ok = tr.Version(channel) + return ok + }, 3*time.Second, 5*time.Millisecond, "no version seeded for channel %q", channel) + return etag, lastModified +} + +func TestFingerprint_StableAndSensitive(t *testing.T) { + sel := mustSelector(t, "flagSetId=fs1") + a := testFlag("fs1", "a", "on") + b := testFlag("fs1", "b", "off") + + // deterministic for identical input + assert.Equal(t, fingerprint(&sel, []model.Flag{a, b}), fingerprint(&sel, []model.Flag{a, b})) + // order independent + assert.Equal(t, fingerprint(&sel, []model.Flag{a, b}), fingerprint(&sel, []model.Flag{b, a})) + + // sensitive to a definition change + aModified := testFlag("fs1", "a", "off") + assert.NotEqual(t, fingerprint(&sel, []model.Flag{a}), fingerprint(&sel, []model.Flag{aModified})) + + // duplicates are not collapsed: a flag appearing twice is a different result set + assert.NotEqual(t, fingerprint(&sel, []model.Flag{a}), fingerprint(&sel, []model.Flag{a, a})) + + // the flagSetId carried on a flag is not part of the identity; the selector the watch was + // opened with is + relabelled := testFlag("some-other-flag-set", "a", "on") + assert.Equal(t, fingerprint(&sel, []model.Flag{a}), fingerprint(&sel, []model.Flag{relabelled})) + + other := mustSelector(t, "source=src1") + assert.NotEqual(t, fingerprint(&sel, []model.Flag{a}), fingerprint(&other, []model.Flag{a})) +} + +// TestTracker_PassesSelectorToStore pins that filtering is delegated to the store by handing it +// the channel's selector, rather than re-implemented by grouping flags in memory. +func TestTracker_PassesSelectorToStore(t *testing.T) { + rec := &recordingStore{} + tr := newTracker(t, rec) + + subscribe(t, tr, "flagSetId=fs1") + + selectors := rec.selectors() + require.Len(t, selectors, 1) + expected := mustSelector(t, "flagSetId=fs1") + assert.Equal(t, expected.ToLogString(), selectors[0].ToLogString()) +} + +func TestTracker_SharesOneWatchPerChannel(t *testing.T) { + rec := &recordingStore{} + tr := newTracker(t, rec) + + subscribe(t, tr, "flagSetId=fs1") + subscribe(t, tr, "flagSetId=fs1") + assert.Len(t, rec.selectors(), 1, "subscribers on one channel must share a store watch") + + subscribe(t, tr, "flagSetId=fs2") + assert.Len(t, rec.selectors(), 2, "a distinct channel needs its own watch") +} + +func TestTracker_Version_NoSubscription(t *testing.T) { + tr := newTracker(t, newTestStore(t)) + + _, _, ok := tr.Version("flagSetId=fs1") + assert.False(t, ok, "no live stream for the channel means no ETag and no 304") +} + +func TestTracker_ReleaseTearsDownLastSubscriberOnly(t *testing.T) { + tr := newTracker(t, newTestStore(t)) + + releaseA, err := tr.Subscribe("flagSetId=fs1", mustSelector(t, "flagSetId=fs1")) + require.NoError(t, err) + releaseB, err := tr.Subscribe("flagSetId=fs1", mustSelector(t, "flagSetId=fs1")) + require.NoError(t, err) + + releaseA() + assert.Len(t, tr.Channels(), 1, "one subscriber leaving must not tear down a shared channel") + + releaseB() + assert.Empty(t, tr.Channels()) + _, _, ok := tr.Version("flagSetId=fs1") + assert.False(t, ok) +} + +func TestTracker_IdenticalUpdateIsSuppressed(t *testing.T) { + s := newTestStore(t) + flags := []model.Flag{testFlag("fs1", "a", "on")} + s.Update("src1", flags, model.Metadata{"flagSetId": "fs1"}, false) + + tr := newTracker(t, s) + subscribe(t, tr, "flagSetId=fs1") + + _, firstModified := versionOf(t, tr, "flagSetId=fs1") + + // Update re-inserts unconditionally, so an identical re-sync still wakes the watch + s.Update("src1", flags, model.Metadata{"flagSetId": "fs1"}, false) + time.Sleep(200 * time.Millisecond) + + _, stillModified := versionOf(t, tr, "flagSetId=fs1") + assert.Equal(t, firstModified, stillModified, "lastModified must not move when nothing changed") +} + +func TestTracker_VersionTracksChanges(t *testing.T) { + s := newTestStore(t) + s.Update("src1", []model.Flag{testFlag("fs1", "a", "on")}, model.Metadata{"flagSetId": "fs1"}, false) + + tr := newTracker(t, s) + subscribe(t, tr, "flagSetId=fs1") + + before, _ := versionOf(t, tr, "flagSetId=fs1") + + s.Update("src1", []model.Flag{testFlag("fs1", "a", "off")}, model.Metadata{"flagSetId": "fs1"}, false) + require.Eventually(t, func() bool { + etag, _, _ := tr.Version("flagSetId=fs1") + return etag != before + }, 3*time.Second, 5*time.Millisecond, "the ETag must move when the config changes") +} + +// closingStore closes the watcher without emitting, simulating store.Watch's error-close path. +type closingStore struct{} + +func (closingStore) Get(context.Context, string, *store.Selector) (model.Flag, model.Metadata, error) { + return model.Flag{}, nil, nil +} + +func (closingStore) GetAll(context.Context, *store.Selector) ([]model.Flag, model.Metadata, error) { + return nil, nil, nil +} + +func (closingStore) Watch(_ context.Context, _ *store.Selector, watcher chan<- store.FlagQueryResult) { + close(watcher) +} + +func (closingStore) Update(string, []model.Flag, model.Metadata, bool) {} + +func TestTracker_StoreErrorInvalidatesVersion(t *testing.T) { + tr := newTracker(t, closingStore{}) + + subscribe(t, tr, "flagSetId=fs1") + + assert.Eventually(t, func() bool { + _, _, ok := tr.Version("flagSetId=fs1") + return !ok + }, 3*time.Second, 5*time.Millisecond, + "a dead watch must stop serving its ETag, or 304s would be stale forever") +} + +func TestTracker_CloseRejectsNewSubscribers(t *testing.T) { + es := eventsource.NewServer() + defer es.Close() + tr := NewTracker(logger.NewLogger(nil, false), newTestStore(t), es) + + release, err := tr.Subscribe("flagSetId=fs1", mustSelector(t, "flagSetId=fs1")) + require.NoError(t, err) + + require.NotPanics(t, tr.Close) + assert.Empty(t, tr.Channels()) + + _, err = tr.Subscribe("flagSetId=fs2", mustSelector(t, "flagSetId=fs2")) + assert.Error(t, err) + + assert.NotPanics(t, release) // releasing after Close is a no-op +} + +// recordingStore records the selector handed to each Watch call, mimicking the real store's +// contract: an immediate initial snapshot on a buffered channel, closed on ctx cancellation. +type recordingStore struct { + mu sync.Mutex + sels []store.Selector +} + +func (r *recordingStore) selectors() []store.Selector { + r.mu.Lock() + defer r.mu.Unlock() + return append([]store.Selector(nil), r.sels...) +} + +func (r *recordingStore) Get(context.Context, string, *store.Selector) (model.Flag, model.Metadata, error) { + return model.Flag{}, nil, nil +} + +func (r *recordingStore) GetAll(context.Context, *store.Selector) ([]model.Flag, model.Metadata, error) { + return nil, nil, nil +} + +func (r *recordingStore) Watch(ctx context.Context, selector *store.Selector, watcher chan<- store.FlagQueryResult) { + r.mu.Lock() + r.sels = append(r.sels, *selector) + r.mu.Unlock() + + go func() { + watcher <- store.FlagQueryResult{} + <-ctx.Done() + close(watcher) + }() +} + +func (r *recordingStore) Update(string, []model.Flag, model.Metadata, bool) {} diff --git a/flagd/pkg/service/flag-evaluation/ofrep/sse_bulk_test.go b/flagd/pkg/service/flag-evaluation/ofrep/sse_bulk_test.go new file mode 100644 index 000000000..c0bb5e091 --- /dev/null +++ b/flagd/pkg/service/flag-evaluation/ofrep/sse_bulk_test.go @@ -0,0 +1,283 @@ +package ofrep + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gorilla/mux" + "github.com/open-feature/flagd/core/pkg/evaluator" + mock "github.com/open-feature/flagd/core/pkg/evaluator/mock" + "github.com/open-feature/flagd/core/pkg/logger" + "github.com/open-feature/flagd/core/pkg/model" + "github.com/open-feature/flagd/core/pkg/service/ofrep" + svc "github.com/open-feature/flagd/flagd/pkg/service" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" +) + +type fakeVersioner struct { + etag string + lastModified int64 + ok bool +} + +func (f fakeVersioner) Version(_ string) (string, int64, bool) { + return f.etag, f.lastModified, f.ok +} + +func newBulkRequest(t *testing.T, selectorHeader, clientEtag string) *http.Request { + t.Helper() + req, err := http.NewRequest(http.MethodPost, "/ofrep/v1/evaluate/flags", bytes.NewReader([]byte{})) + require.NoError(t, err) + req.Host = "flagd.example:8016" + if selectorHeader != "" { + req.Header.Set(svc.FLAGD_SELECTOR_HEADER, selectorHeader) + } + if clientEtag != "" { + req.Header.Set("If-None-Match", clientEtag) + } + return req +} + +func serveBulk(h handler, req *http.Request) *httptest.ResponseRecorder { + recorder := httptest.NewRecorder() + router := mux.NewRouter() + router.HandleFunc(bulkEvaluation, h.HandleBulkEvaluation) + router.ServeHTTP(recorder, req) + return recorder +} + +func TestHandleBulkEvaluation_NotModified(t *testing.T) { + log := logger.NewLogger(nil, false) + eval := mock.NewMockIEvaluator(gomock.NewController(t)) + // evaluation must be skipped on a 304 + eval.EXPECT().ResolveAllValues(gomock.Any(), gomock.Any(), gomock.Any()).Times(0) + + h := handler{ + Logger: log, + evaluator: eval, + versioner: fakeVersioner{etag: "abc123", ok: true}, + sseEnabled: true, + sseInactivityDelaySec: 120, + } + + recorder := serveBulk(h, newBulkRequest(t, "flagSetId=fs1", `"abc123"`)) + + assert.Equal(t, http.StatusNotModified, recorder.Code) + assert.Equal(t, `"abc123"`, recorder.Header().Get("ETag")) +} + +// The 304 decision must use If-None-Match (the client's cached version), not the ADR-0008 +// flagConfigEtag change-trigger metadata. +func TestHandleBulkEvaluation_ConditionalUsesIfNoneMatch(t *testing.T) { + log := logger.NewLogger(nil, false) + const current = "etag-v2" + + tests := []struct { + name string + flagConfigEtag string // query param (change-trigger metadata) + ifNoneMatch string // header (cache validator) + wantStatus int + wantEval bool + }{ + { + name: "stale cache echoing new trigger etag is served fresh (regression)", + flagConfigEtag: "etag-v2", + ifNoneMatch: `"etag-v1"`, + wantStatus: http.StatusOK, + wantEval: true, + }, + { + name: "current cache validator yields 304", + ifNoneMatch: `"etag-v2"`, + wantStatus: http.StatusNotModified, + wantEval: false, + }, + { + name: "trigger etag alone (no validator) is served fresh", + flagConfigEtag: "etag-v2", + wantStatus: http.StatusOK, + wantEval: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + eval := mock.NewMockIEvaluator(gomock.NewController(t)) + expect := eval.EXPECT().ResolveAllValues(gomock.Any(), gomock.Any(), gomock.Any()). + Return([]evaluator.AnyValue{}, model.Metadata{}, nil) + if tt.wantEval { + expect.Times(1) + } else { + expect.Times(0) + } + + h := handler{ + Logger: log, + evaluator: eval, + versioner: fakeVersioner{etag: current, ok: true}, + sseEnabled: true, + } + + url := "/ofrep/v1/evaluate/flags" + if tt.flagConfigEtag != "" { + url += "?flagConfigEtag=" + tt.flagConfigEtag + } + req, err := http.NewRequest(http.MethodPost, url, bytes.NewReader([]byte{})) + require.NoError(t, err) + req.Header.Set(svc.FLAGD_SELECTOR_HEADER, "flagSetId=fs1") + if tt.ifNoneMatch != "" { + req.Header.Set("If-None-Match", tt.ifNoneMatch) + } + + recorder := serveBulk(h, req) + + assert.Equal(t, tt.wantStatus, recorder.Code) + assert.Equal(t, `"`+current+`"`, recorder.Header().Get("ETag")) + }) + } +} + +func TestHandleBulkEvaluation_AdvertisesEventStreams(t *testing.T) { + log := logger.NewLogger(nil, false) + + tests := []struct { + name string + publicURL string + selectorHeader string + wantOrigin string // expected endpoint.origin ("" => omitted) + wantRequestURI string + }{ + { + name: "flagSetId selector, origin omitted", + selectorHeader: "flagSetId=fs1", + wantOrigin: "", + wantRequestURI: "/ofrep/v1/sse/flagSetId=fs1", + }, + { + // previously fell back to the catch-all channel, and so was woken by every change + name: "source selector gets its own channel", + selectorHeader: "source=src1", + wantOrigin: "", + wantRequestURI: "/ofrep/v1/sse/source=src1", + }, + { + // advertised verbatim, so the SSE endpoint parses it back to the same selector + name: "bare selector expression is advertised verbatim", + selectorHeader: "mySource", + wantOrigin: "", + wantRequestURI: "/ofrep/v1/sse/mySource", + }, + { + // the channel is one path segment, so a source path must be escaped into it + name: "source containing a slash is path-escaped", + selectorHeader: "source=./mySource", + wantOrigin: "", + wantRequestURI: "/ofrep/v1/sse/source=.%2FmySource", + }, + { + // "flags with no flagSetId"; the internal nil flagSetId must never be advertised + name: "empty flagSetId selector", + selectorHeader: "flagSetId=", + wantOrigin: "", + wantRequestURI: "/ofrep/v1/sse/flagSetId=", + }, + { + name: "no selector (catch-all), origin omitted", + selectorHeader: "", + wantOrigin: "", + wantRequestURI: "/ofrep/v1/sse", + }, + { + name: "public url sets origin", + publicURL: "https://flags.example.com/", + selectorHeader: "flagSetId=fs1", + wantOrigin: "https://flags.example.com", + wantRequestURI: "/ofrep/v1/sse/flagSetId=fs1", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + eval := mock.NewMockIEvaluator(gomock.NewController(t)) + eval.EXPECT().ResolveAllValues(gomock.Any(), gomock.Any(), gomock.Any()). + Return([]evaluator.AnyValue{}, model.Metadata{}, nil) + + h := handler{ + Logger: log, + evaluator: eval, + versioner: fakeVersioner{etag: "etag-1", ok: true}, + sseEnabled: true, + sseInactivityDelaySec: 120, + ssePublicURL: tt.publicURL, + } + + recorder := serveBulk(h, newBulkRequest(t, tt.selectorHeader, "")) + + require.Equal(t, http.StatusOK, recorder.Code) + assert.Equal(t, `"etag-1"`, recorder.Header().Get("ETag")) + + var resp ofrep.BulkEvaluationResponse + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &resp)) + require.Len(t, resp.EventStreams, 1) + + es := resp.EventStreams[0] + assert.Equal(t, "sse", es.Type) + assert.Equal(t, 120, es.InactivityDelaySec) + require.NotNil(t, es.Endpoint) + assert.Equal(t, tt.wantOrigin, es.Endpoint.Origin) + assert.Equal(t, tt.wantRequestURI, es.Endpoint.RequestUri) + }) + } +} + +func TestHandleBulkEvaluation_SSEDisabled_NoEventStreams(t *testing.T) { + log := logger.NewLogger(nil, false) + eval := mock.NewMockIEvaluator(gomock.NewController(t)) + eval.EXPECT().ResolveAllValues(gomock.Any(), gomock.Any(), gomock.Any()). + Return([]evaluator.AnyValue{}, model.Metadata{}, nil) + + // sseEnabled false, no versioner: behaves like the legacy handler + h := handler{Logger: log, evaluator: eval} + + recorder := serveBulk(h, newBulkRequest(t, "", "")) + + require.Equal(t, http.StatusOK, recorder.Code) + assert.Empty(t, recorder.Header().Get("ETag")) + + var resp ofrep.BulkEvaluationResponse + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &resp)) + assert.Empty(t, resp.EventStreams) +} + +// TestHandleBulkEvaluation_NoLiveStream_NoETag pins the contract that follows from tracking +// versions per live subscription: with no stream open there is no cache validator to offer, so +// the response is a plain 200 that still advertises eventStreams. +func TestHandleBulkEvaluation_NoLiveStream_NoETag(t *testing.T) { + log := logger.NewLogger(nil, false) + eval := mock.NewMockIEvaluator(gomock.NewController(t)) + eval.EXPECT().ResolveAllValues(gomock.Any(), gomock.Any(), gomock.Any()). + Return([]evaluator.AnyValue{}, model.Metadata{}, nil) + + h := handler{ + Logger: log, + evaluator: eval, + versioner: fakeVersioner{ok: false}, + sseEnabled: true, + sseInactivityDelaySec: 120, + } + + recorder := serveBulk(h, newBulkRequest(t, "flagSetId=fs1", `"stale-etag"`)) + + require.Equal(t, http.StatusOK, recorder.Code, "without a tracked version we must not serve a 304") + assert.Empty(t, recorder.Header().Get("ETag")) + + var resp ofrep.BulkEvaluationResponse + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &resp)) + require.Len(t, resp.EventStreams, 1) + assert.NotContains(t, resp.Metadata, "flagConfigLastModified") +}