diff --git a/core/pkg/store/store.go b/core/pkg/store/store.go index ee25d9910..a0bfbdd2f 100644 --- a/core/pkg/store/store.go +++ b/core/pkg/store/store.go @@ -206,6 +206,18 @@ func (s *Store) GetAll(ctx context.Context, selector *Selector) ([]model.Flag, m return flags, queryMeta, nil } +// watchSelector returns a channel that will be closed when the flags matching the given selector are modified. +func (s *Store) WatchSelector(selector *Selector) <-chan struct{} { + it, err := s.selectOrAll(selector) + if err != nil { + // return a closed channel on error + ch := make(chan struct{}) + close(ch) + return ch + } + return it.WatchCh() +} + type flagIdentifier struct { flagSetId string key string diff --git a/docs/reference/flagd-cli/flagd_start.md b/docs/reference/flagd-cli/flagd_start.md index 1a853f956..6c6ed5e39 100644 --- a/docs/reference/flagd-cli/flagd_start.md +++ b/docs/reference/flagd-cli/flagd_start.md @@ -19,6 +19,7 @@ flagd start [flags] -z, --log-format string Set the logging format, e.g. console or json (default "console") -m, --management-port int32 Port for management operations (default 8014) -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 + --ofrep-cache-capacity int32 Max number of selectors to cache for OFREP bulk evaluation ETags (0 = unlimited) (default 100) -r, --ofrep-port int32 ofrep service port (default 8016) -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 diff --git a/docs/reference/flagd-ofrep.md b/docs/reference/flagd-ofrep.md index f28ac3b85..4f7662b0f 100644 --- a/docs/reference/flagd-ofrep.md +++ b/docs/reference/flagd-ofrep.md @@ -23,4 +23,10 @@ To evaluate all flags currently configured at flagd, use OFREP bulk evaluation r curl -X POST 'http://localhost:8016/ofrep/v1/evaluate/flags' ``` +## Evaluation Caching + +The bulk evaluation endpoint caches responses per selector to avoid redundant evaluations. Clients can use the `If-None-Match` header with a previously received `ETag` to check if the cache is still valid. When the ETag matches, flagd returns the cached response without re-evaluating. + +**Important**: The ETag only corresponds the flag configuration version, not the evaluation context. Clients must not send a cached ETag when their evaluation context has changed, otherwise they may receive stale results. + See the [cheat sheet](./cheat-sheet.md#ofrep-api-http) for more OFREP examples including context-sensitive evaluation and selectors. diff --git a/flagd/cmd/start.go b/flagd/cmd/start.go index 83745dd5a..33d808f44 100644 --- a/flagd/cmd/start.go +++ b/flagd/cmd/start.go @@ -23,6 +23,7 @@ const ( managementPortFlagName = "management-port" metricsExporter = "metrics-exporter" ofrepPortFlagName = "ofrep-port" + ofrepCacheCapacityFlagName = "ofrep-cache-capacity" otelCollectorURI = "otel-collector-uri" otelCertPathFlagName = "otel-cert-path" otelKeyPathFlagName = "otel-key-path" @@ -52,6 +53,8 @@ func init() { flags.Int32P(portFlagName, "p", 8013, "Port to listen on") flags.Int32P(syncPortFlagName, "g", 8015, "gRPC Sync port") flags.Int32P(ofrepPortFlagName, "r", 8016, "ofrep service port") + flags.Int32(ofrepCacheCapacityFlagName, 100, + "Max number of selectors to cache for OFREP bulk evaluation ETags (0 = unlimited)") flags.StringP(socketPathFlagName, "d", "", "Flagd unix socket path. "+ "With grpc the evaluations service will become available on this address. "+ @@ -113,6 +116,7 @@ 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(ofrepCacheCapacityFlagName, flags.Lookup(ofrepCacheCapacityFlagName)) _ = viper.BindPFlag(contextValueFlagName, flags.Lookup(contextValueFlagName)) _ = viper.BindPFlag(headerToContextKeyFlagName, flags.Lookup(headerToContextKeyFlagName)) _ = viper.BindPFlag(streamDeadlineFlagName, flags.Lookup(streamDeadlineFlagName)) @@ -177,6 +181,7 @@ var startCmd = &cobra.Command{ MetricExporter: viper.GetString(metricsExporter), ManagementPort: viper.GetUint16(managementPortFlagName), OfrepServicePort: viper.GetUint16(ofrepPortFlagName), + OfrepCacheCapacity: viper.GetInt(ofrepCacheCapacityFlagName), OtelCollectorURI: viper.GetString(otelCollectorURI), OtelCertPath: viper.GetString(otelCertPathFlagName), OtelKeyPath: viper.GetString(otelKeyPathFlagName), diff --git a/flagd/pkg/runtime/from_config.go b/flagd/pkg/runtime/from_config.go index 08fcc5b0c..f75decdaf 100644 --- a/flagd/pkg/runtime/from_config.go +++ b/flagd/pkg/runtime/from_config.go @@ -27,6 +27,7 @@ type Config struct { MetricExporter string ManagementPort uint16 OfrepServicePort uint16 + OfrepCacheCapacity int OtelCollectorURI string OtelCertPath string OtelKeyPath string @@ -104,10 +105,11 @@ func FromConfig(logger *logger.Logger, version string, config Config) (*Runtime, recorder) // ofrep service - ofrepService, err := ofrep.NewOfrepService(jsonEvaluator, config.CORS, ofrep.SvcConfiguration{ - Logger: logger.WithFields(zap.String("component", "OFREPService")), - Port: config.OfrepServicePort, - ServiceName: svcName, + ofrepService, err := ofrep.NewOfrepService(jsonEvaluator, store, config.CORS, ofrep.SvcConfiguration{ + Logger: logger.WithFields(zap.String("component", "OFREPService")), + Port: config.OfrepServicePort, + CacheCapacity: config.OfrepCacheCapacity, + ServiceName: svcName, MetricsRecorder: recorder, }, config.ContextValues, diff --git a/flagd/pkg/service/flag-evaluation/ofrep/handler.go b/flagd/pkg/service/flag-evaluation/ofrep/handler.go index cbc8d435b..f12cc1cdc 100644 --- a/flagd/pkg/service/flag-evaluation/ofrep/handler.go +++ b/flagd/pkg/service/flag-evaluation/ofrep/handler.go @@ -19,12 +19,26 @@ import ( "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" "go.opentelemetry.io/otel" "go.opentelemetry.io/otel/trace" + "go.uber.org/zap" ) +// ISelectorVersionTracker defines the interface for selector version tracking. +// it enables ETag-based cache validation for bulk evaluations. +type ISelectorVersionTracker interface { + // ETag returns the current ETag for a selector, or empty if not tracked + ETag(selectorExpression string) string + // Track starts tracking a selector and returns its initial ETag + Track(selectorExpression string) string +} + const ( - key = "key" - singleEvaluation = "/ofrep/v1/evaluate/flags/{key}" - bulkEvaluation = "/ofrep/v1/evaluate/{path:flags\\/|flags}" + key = "key" + singleEvaluation = "/ofrep/v1/evaluate/flags/{key}" + bulkEvaluation = "/ofrep/v1/evaluate/{path:flags\\/|flags}" + headerETag = "ETag" + headerIfNoneMatch = "If-None-Match" + headerContentType = "Content-Type" + contentTypeJSON = "application/json" ) type handler struct { @@ -33,6 +47,7 @@ type handler struct { contextValues map[string]any headerToContextKeyMappings map[string]string tracer trace.Tracer + versionTracker ISelectorVersionTracker } func NewOfrepHandler( @@ -42,6 +57,7 @@ func NewOfrepHandler( headerToContextKeyMappings map[string]string, metricsRecorder telemetry.IMetricsRecorder, serviceName string, + versionTracker ISelectorVersionTracker, ) http.Handler { h := handler{ Logger: logger, @@ -49,6 +65,7 @@ func NewOfrepHandler( contextValues: contextValues, headerToContextKeyMappings: headerToContextKeyMappings, tracer: otel.Tracer("flagd.ofrep.v1"), + versionTracker: versionTracker, } router := mux.NewRouter() @@ -117,6 +134,19 @@ func (h *handler) HandleBulkEvaluation(w http.ResponseWriter, r *http.Request) { evaluationContext := flagdContext(h.Logger, requestID, request, h.contextValues, r.Header, h.headerToContextKeyMappings) selectorExpression := r.Header.Get(service.FLAGD_SELECTOR_HEADER) + + // check if client's ETag matches current version - return 304 if unchanged + ifNoneMatch := r.Header.Get(headerIfNoneMatch) + if h.versionTracker != nil && ifNoneMatch != "" { + currentETag := h.versionTracker.ETag(selectorExpression) + if currentETag != "" && ifNoneMatch == currentETag { +h.Logger.Debug("ETag match, returning 304", zap.String("selector", selectorExpression)) + w.Header().Add(headerETag, currentETag) + w.WriteHeader(http.StatusNotModified) + return + } + } + selector := store.NewSelector(selectorExpression) ctx := context.WithValue(r.Context(), store.SelectorContextKey{}, selector) @@ -128,7 +158,36 @@ func (h *handler) HandleBulkEvaluation(w http.ResponseWriter, r *http.Request) { 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) + response := ofrep.BulkEvaluationResponseFrom(evaluations, metadata) + h.writeBulkEvaluationResponse(w, r, selectorExpression, response) + } +} + +// writes the bulk evaluation response with ETag support +func (h *handler) writeBulkEvaluationResponse(w http.ResponseWriter, _ *http.Request, selectorExpression string, response ofrep.BulkEvaluationResponse) { + // marshal the response + body, err := json.Marshal(response) + if err != nil { + h.Logger.Warn("error marshalling response", zap.Error(err)) + w.WriteHeader(http.StatusInternalServerError) + return + } + + // track this selector and get ETag for response + var eTag string + if h.versionTracker != nil { + eTag = h.versionTracker.Track(selectorExpression) + } + + // write response with ETag + w.Header().Add(headerContentType, contentTypeJSON) + if eTag != "" { + w.Header().Add(headerETag, eTag) + } + w.WriteHeader(http.StatusOK) + _, err = w.Write(body) + if err != nil { + h.Logger.Warn("error while writing response", zap.Error(err)) } } @@ -137,16 +196,16 @@ func (h *handler) writeJSONToResponse(status int, payload interface{}, w http.Re marshal, err := json.Marshal(payload) if err != nil { // always a 500 - h.Logger.Warn(fmt.Sprintf("error marshelling the response: %v", err)) + h.Logger.Warn("error marshalling the response", zap.Error(err)) w.WriteHeader(http.StatusInternalServerError) return } - w.Header().Add("Content-Type", "application/json") + w.Header().Add(headerContentType, contentTypeJSON) w.WriteHeader(status) _, err = w.Write(marshal) if err != nil { - h.Logger.Warn(fmt.Sprintf("error while writing response: %v", err)) + h.Logger.Warn("error while writing response", zap.Error(err)) } } diff --git a/flagd/pkg/service/flag-evaluation/ofrep/handler_test.go b/flagd/pkg/service/flag-evaluation/ofrep/handler_test.go index 2ef114d41..522a0af35 100644 --- a/flagd/pkg/service/flag-evaluation/ofrep/handler_test.go +++ b/flagd/pkg/service/flag-evaluation/ofrep/handler_test.go @@ -2,6 +2,7 @@ package ofrep import ( "bytes" + "context" "encoding/json" "errors" "io" @@ -9,6 +10,7 @@ import ( "net/http/httptest" "reflect" "testing" + "time" "github.com/gorilla/mux" "github.com/open-feature/flagd/core/pkg/evaluator" @@ -16,9 +18,31 @@ import ( "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" + "github.com/open-feature/flagd/core/pkg/store" "go.uber.org/mock/gomock" ) +// testFlagStore is a mock FlagStore for handler tests +type testFlagStore struct { + flags []model.Flag + watchCh chan struct{} +} + +func newTestFlagStore() *testFlagStore { + return &testFlagStore{ + flags: []model.Flag{{Key: "test-flag", State: "ENABLED"}}, + watchCh: make(chan struct{}), + } +} + +func (m *testFlagStore) GetAll(_ context.Context, _ *store.Selector) ([]model.Flag, model.Metadata, error) { + return m.flags, model.Metadata{}, nil +} + +func (m *testFlagStore) WatchSelector(_ *store.Selector) <-chan struct{} { + return m.watchCh +} + var flagKey = "key" var successValue = evaluator.AnyValue{ @@ -281,7 +305,7 @@ func TestWriteJSONResponse(t *testing.T) { var rsp evaluator.AnyValue err = json.Unmarshal(b, &rsp) if err != nil { - t.Errorf("error unmarshelling body: %v", err) + t.Errorf("error unmarshaling body: %v", err) } if !reflect.DeepEqual(test.payload, rsp) { @@ -290,3 +314,226 @@ func TestWriteJSONResponse(t *testing.T) { }) } } + +func TestWriteBulkEvaluationResponse_ETag(t *testing.T) { + log := logger.NewLogger(nil, false) + + // create version tracker with mock store + mockStore := newTestFlagStore() + tracker := NewSelectorVersionTracker(log, mockStore, 0) + defer tracker.Close() + + h := handler{Logger: log, versionTracker: tracker} + + // test response + response := ofrep.BulkEvaluationResponse{ + Flags: []interface{}{ + ofrep.EvaluationSuccess{ + Key: "test-flag", + Value: true, + Reason: model.StaticReason, + Variant: "on", + }, + }, + Metadata: model.Metadata{}, + } + + selectorExpression := "flagSetId=test-set" + + tests := []struct { + name string + ifNoneMatch string + expectedStatus int + expectedHasETag bool + expectedHasBody bool + }{ + { + name: "no If-None-Match header returns 200 with body and ETag", + ifNoneMatch: "", + expectedStatus: http.StatusOK, + expectedHasETag: true, + expectedHasBody: true, + }, + { + name: "non-matching If-None-Match header returns 200 with body and ETag", + ifNoneMatch: "\"some-invalid-etag-lmao\"", + expectedStatus: http.StatusOK, + expectedHasETag: true, + expectedHasBody: true, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + request := httptest.NewRequest(http.MethodPost, "/ofrep/v1/evaluate/flags", bytes.NewReader([]byte{})) + if test.ifNoneMatch != "" { + request.Header.Set("If-None-Match", test.ifNoneMatch) + } + + recorder := httptest.NewRecorder() + h.writeBulkEvaluationResponse(recorder, request, selectorExpression, response) + + if test.expectedStatus != recorder.Code { + t.Errorf("expected status code %d, but got %d", test.expectedStatus, recorder.Code) + } + + eTagHeader := recorder.Header().Get("ETag") + if test.expectedHasETag && eTagHeader == "" { + t.Error("expected ETag header to be present, but it was missing") + } + if !test.expectedHasETag && eTagHeader != "" { + t.Error("expected ETag header to be absent, but it was present") + } + + body := recorder.Body.String() + if test.expectedHasBody && body == "" { + t.Error("expected response body, but got empty body") + } + if !test.expectedHasBody && body != "" { + t.Errorf("expected no response body, but got: %s", body) + } + + // for 200 responses, verify Content-Type header is present + if test.expectedStatus == http.StatusOK && recorder.Header().Get("Content-Type") != "application/json" { + t.Error("expected Content-Type header to be application/json, but it was missing") + } + }) + } + + // test matching If-None-Match - the ETag should match since Track is idempotent + t.Run("matching If-None-Match header still returns 200 (no 304 in writeBulkEvaluationResponse)", func(t *testing.T) { + // get the current ETag for this selector + currentETag := tracker.ETag(selectorExpression) + + request := httptest.NewRequest(http.MethodPost, "/ofrep/v1/evaluate/flags", bytes.NewReader([]byte{})) + request.Header.Set("If-None-Match", currentETag) + + recorder := httptest.NewRecorder() + h.writeBulkEvaluationResponse(recorder, request, selectorExpression, response) + }) +} + +func TestHandleBulkEvaluation_304NotModified(t *testing.T) { + log := logger.NewLogger(nil, false) + + mockStore := newTestFlagStore() + tracker := NewSelectorVersionTracker(log, mockStore, 0) + defer tracker.Close() + + // pre-track the empty selector + selectorExpression := "" + cachedETag := tracker.Track(selectorExpression) + + // create handler with version tracker - evaluator should NOT be called for 304 + eval := mock.NewMockIEvaluator(gomock.NewController(t)) + + h := handler{Logger: log, evaluator: eval, versionTracker: tracker} + + // request WITH matching If-None-Match header - should return 304 + request := httptest.NewRequest(http.MethodPost, "/ofrep/v1/evaluate/flags", bytes.NewReader([]byte{})) + request.Header.Set("If-None-Match", cachedETag) + recorder := httptest.NewRecorder() + + router := mux.NewRouter() + router.HandleFunc(bulkEvaluation, h.HandleBulkEvaluation) + router.ServeHTTP(recorder, request) + + if recorder.Code != http.StatusNotModified { + t.Errorf("expected status 304, got %d", recorder.Code) + } + + if recorder.Header().Get("ETag") == "" { + t.Error("expected ETag header to be present") + } +} + +func TestHandleBulkEvaluation_NewSelector(t *testing.T) { + log := logger.NewLogger(nil, false) + + // create version tracker with mock store + mockStore := newTestFlagStore() + tracker := NewSelectorVersionTracker(log, mockStore, 0) + defer tracker.Close() + + // create handler with version tracker + eval := mock.NewMockIEvaluator(gomock.NewController(t)) + eval.EXPECT().ResolveAllValues(gomock.Any(), gomock.Any(), gomock.Any()). + Return([]evaluator.AnyValue{successValue}, model.Metadata{}, nil) + + h := handler{Logger: log, evaluator: eval, versionTracker: tracker} + + request := httptest.NewRequest(http.MethodPost, "/ofrep/v1/evaluate/flags", bytes.NewReader([]byte{})) + recorder := httptest.NewRecorder() + + router := mux.NewRouter() + router.HandleFunc(bulkEvaluation, h.HandleBulkEvaluation) + router.ServeHTTP(recorder, request) + + if recorder.Code != http.StatusOK { + t.Errorf("expected status 200, got %d", recorder.Code) + } + + if recorder.Header().Get("ETag") == "" { + t.Error("expected ETag header to be present") + } + + etag := tracker.ETag("") + if etag == "" { + t.Error("expected selector to be tracked") + } +} + +func TestVersionBumpOnStoreUpdate(t *testing.T) { + log := logger.NewLogger(nil, false) + + // create a real store + flagStore, err := store.NewStore(log, []string{"test-source"}) + if err != nil { + t.Fatalf("failed to create store: %v", err) + } + + // add initial flags + flagStore.Update("test-source", []model.Flag{ + { + Key: "test-flag", + State: "ENABLED", + DefaultVariant: "on", + Variants: map[string]any{"on": true, "off": false}, + }, + }, model.Metadata{}) + + // create version tracker with watch provider (the store) + tracker := NewSelectorVersionTracker(log, flagStore, 0) + defer tracker.Close() + + // track the empty selector (matches all flags) + etagBefore := tracker.Track("") + + if etagBefore == "" { + t.Fatal("expected non-empty ETag after tracking") + } + + // update the store with different flag content + flagStore.Update("test-source", []model.Flag{ + { + Key: "test-flag", + State: "ENABLED", + DefaultVariant: "off", // changed! + Variants: map[string]any{"on": true, "off": false}, + }, + }, model.Metadata{}) + + // give time for watch goroutine to process the update + time.Sleep(50 * time.Millisecond) + + // ETag should have changed because content changed + etagAfter := tracker.ETag("") + + if etagAfter == "" { + t.Fatal("expected non-empty ETag after update") + } + + if etagBefore == etagAfter { + t.Errorf("expected ETag to change after store update, but got same value: %s", etagBefore) + } +} diff --git a/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service.go b/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service.go index bc5689ee4..27e369c82 100644 --- a/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service.go +++ b/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service.go @@ -9,6 +9,7 @@ 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/rs/cors" "golang.org/x/sync/errgroup" @@ -22,19 +23,30 @@ type IOfrepService interface { type SvcConfiguration struct { Logger *logger.Logger Port uint16 + CacheCapacity int ServiceName string MetricsRecorder telemetry.IMetricsRecorder } type Service struct { - logger *logger.Logger - port uint16 - server *http.Server + logger *logger.Logger + port uint16 + server *http.Server + versionTracker *SelectorVersionTracker } func NewOfrepService( - evaluator evaluator.IEvaluator, origins []string, cfg SvcConfiguration, contextValues map[string]any, headerToContextKeyMappings map[string]string, + evaluator evaluator.IEvaluator, + flagStore *store.Store, + origins []string, + cfg SvcConfiguration, + contextValues map[string]any, + headerToContextKeyMappings map[string]string, ) (*Service, error) { + // create the version tracker with watch-based invalidation + // the store implements WatchProvider interface for targeted invalidation + versionTracker := NewSelectorVersionTracker(cfg.Logger, flagStore, cfg.CacheCapacity) + corsMW := cors.New(cors.Options{ AllowedOrigins: origins, AllowedMethods: []string{http.MethodPost}, @@ -47,6 +59,7 @@ func NewOfrepService( headerToContextKeyMappings, cfg.MetricsRecorder, cfg.ServiceName, + versionTracker, )) server := http.Server{ @@ -56,9 +69,10 @@ func NewOfrepService( } return &Service{ - logger: cfg.Logger, - port: cfg.Port, - server: &server, + logger: cfg.Logger, + port: cfg.Port, + server: &server, + versionTracker: versionTracker, }, nil } @@ -78,6 +92,12 @@ func (s Service) Start(ctx context.Context) error { group.Go(func() error { <-gCtx.Done() s.logger.Info("shutting down ofrep service") + + // close the version tracker to stop watch goroutines + if s.versionTracker != nil { + s.versionTracker.Close() + } + err := s.server.Close() if err != nil { return fmt.Errorf("error from ofrep server shutdown: %w", err) 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 ee30780f9..cd5fa0546 100644 --- a/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service_test.go +++ b/flagd/pkg/service/flag-evaluation/ofrep/ofrep_service_test.go @@ -12,6 +12,7 @@ import ( 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/store" "github.com/open-feature/flagd/core/pkg/telemetry" "go.uber.org/mock/gomock" "golang.org/x/sync/errgroup" @@ -24,14 +25,20 @@ func Test_OfrepServiceStartStop(t *testing.T) { eval.EXPECT().ResolveAllValues(gomock.Any(), gomock.Any(), gomock.Any()). Return([]evaluator.AnyValue{}, model.Metadata{}, nil) + log := logger.NewLogger(nil, false) + flagStore, err := store.NewStore(log, []string{}) + if err != nil { + t.Fatalf("error creating store: %v", err) + } + cfg := SvcConfiguration{ - Logger: logger.NewLogger(nil, false), - Port: uint16(port), + Logger: log, + Port: uint16(port), ServiceName: "test-service", MetricsRecorder: &telemetry.NoopMetricsRecorder{}, } - service, err := NewOfrepService(eval, []string{"*"}, cfg, nil, nil) + service, err := NewOfrepService(eval, flagStore, []string{"*"}, cfg, nil, nil) if err != nil { t.Fatalf("error creating the ofrep service: %v", err) } diff --git a/flagd/pkg/service/flag-evaluation/ofrep/selector_version_tracker.go b/flagd/pkg/service/flag-evaluation/ofrep/selector_version_tracker.go new file mode 100644 index 000000000..8d8a5756f --- /dev/null +++ b/flagd/pkg/service/flag-evaluation/ofrep/selector_version_tracker.go @@ -0,0 +1,207 @@ +package ofrep + +import ( + "context" + "crypto/md5" + "encoding/json" + "fmt" + "sync" + + "github.com/open-feature/flagd/core/pkg/logger" + "github.com/open-feature/flagd/core/pkg/model" + "github.com/open-feature/flagd/core/pkg/store" +) + +// FlagStore is an interface for querying flags and watching for changes. +type FlagStore interface { + GetAll(ctx context.Context, selector *store.Selector) ([]model.Flag, model.Metadata, error) + WatchSelector(selector *store.Selector) <-chan struct{} +} + +// trackedSelector holds the state for a single tracked selector +type trackedSelector struct { + etag string + cancel context.CancelFunc +} + +// SelectorVersionTracker tracks content hashes for selectors to enable ETag-based caching. +type SelectorVersionTracker struct { + logger *logger.Logger + flagStore FlagStore + mu sync.RWMutex + selectors map[string]*trackedSelector // single map for all per-selector state + insertOrder []string // FIFO order for eviction + maxCapacity int + ctx context.Context + cancel context.CancelFunc +} + +// NewSelectorVersionTracker creates a new version tracker with watch-based invalidation. +// maxCapacity limits the number of tracked selectors (0 = unlimited). +func NewSelectorVersionTracker(logger *logger.Logger, flagStore FlagStore, maxCapacity int) *SelectorVersionTracker { + ctx, cancel := context.WithCancel(context.Background()) + return &SelectorVersionTracker{ + logger: logger, + flagStore: flagStore, + selectors: make(map[string]*trackedSelector), + insertOrder: make([]string, 0), + maxCapacity: maxCapacity, + ctx: ctx, + cancel: cancel, + } +} + +// Close shuts down the tracker and stops all watch goroutines +func (t *SelectorVersionTracker) Close() { + t.cancel() +} + +// ETag returns the current ETag for a selector. +// returns empty string if the selector has never been tracked. +func (t *SelectorVersionTracker) ETag(selectorExpression string) string { + t.mu.RLock() + defer t.mu.RUnlock() + if s, ok := t.selectors[selectorExpression]; ok { + return s.etag + } + return "" +} + +// Track starts tracking a selector and returns its current content-based ETag. +// if already tracking, returns the cached ETag without recomputing. +func (t *SelectorVersionTracker) Track(selectorExpression string) string { + t.mu.Lock() + defer t.mu.Unlock() + + // if already tracking, return cached ETag + if s, exists := t.selectors[selectorExpression]; exists { + return s.etag + } + + // evict oldest if at capacity + if t.maxCapacity > 0 && len(t.selectors) >= t.maxCapacity { + t.evictOldest() + } + + // compute content-based ETag + etag := t.computeETag(selectorExpression) + + // start watching for changes + var watchCancel context.CancelFunc + if t.flagStore != nil { + selector := store.NewSelector(selectorExpression) + watchCh := t.flagStore.WatchSelector(&selector) + var watchCtx context.Context + watchCtx, watchCancel = context.WithCancel(t.ctx) + go t.watchAndRecompute(watchCtx, selectorExpression, watchCh) + } + + t.selectors[selectorExpression] = &trackedSelector{etag: etag, cancel: watchCancel} + t.insertOrder = append(t.insertOrder, selectorExpression) + + t.logger.Debug(fmt.Sprintf("tracking selector '%s' with ETag %s", selectorExpression, etag)) + return etag +} + +// computeETag generates a content-based ETag by hashing the flags for a selector. +// this ensures ETags are consistent across replicas with the same flag content. +func (t *SelectorVersionTracker) computeETag(selectorExpression string) string { + if t.flagStore == nil { + return "" + } + + selector := store.NewSelector(selectorExpression) + flags, metadata, err := t.flagStore.GetAll(t.ctx, &selector) + if err != nil { + t.logger.Warn(fmt.Sprintf("error getting flags for selector '%s': %v", selectorExpression, err)) + return "" + } + + // create a hashable representation that includes the key + // (model.Flag.MarshalJSON omits the key, so we need to include it explicitly) + type flagForHash struct { + Key string + State string + DefaultVariant string + Variants map[string]any + Targeting json.RawMessage + Metadata model.Metadata + } + hashableFlags := make([]flagForHash, len(flags)) + for i, f := range flags { + hashableFlags[i] = flagForHash{ + Key: f.Key, + State: f.State, + DefaultVariant: f.DefaultVariant, + Variants: f.Variants, + Targeting: f.Targeting, + Metadata: f.Metadata, + } + } + + // serialize flags and metadata to create deterministic hash + data, err := json.Marshal(struct { + Flags []flagForHash + Metadata model.Metadata + }{Flags: hashableFlags, Metadata: metadata}) + if err != nil { + t.logger.Warn(fmt.Sprintf("error marshaling flags for selector '%s': %v", selectorExpression, err)) + return "" + } + + hash := md5.Sum(data) + return fmt.Sprintf("\"%x\"", hash) +} + +// evictOldest removes the oldest tracked selector (FIFO). +// caller must hold the mutex. +func (t *SelectorVersionTracker) evictOldest() { + if len(t.insertOrder) == 0 { + return + } + oldest := t.insertOrder[0] + t.insertOrder = t.insertOrder[1:] + if s, ok := t.selectors[oldest]; ok { + if s.cancel != nil { + s.cancel() + } + delete(t.selectors, oldest) + } + t.logger.Warn(fmt.Sprintf("evicted selector '%s' from version tracker, consider increasing ofrep-cache-capacity", oldest)) +} + +// watchAndRecompute monitors a watch channel and recomputes the ETag when flags change +func (t *SelectorVersionTracker) watchAndRecompute(ctx context.Context, selectorExpression string, watchCh <-chan struct{}) { + for { + select { + case <-ctx.Done(): + return + case <-watchCh: + t.recomputeETag(selectorExpression) + + // re-establish watch for future changes + if t.flagStore != nil { + selector := store.NewSelector(selectorExpression) + watchCh = t.flagStore.WatchSelector(&selector) + } else { + return + } + } + } +} + +// recomputeETag recomputes and updates the content-based ETag for a selector +func (t *SelectorVersionTracker) recomputeETag(selectorExpression string) { + etag := t.computeETag(selectorExpression) + + t.mu.Lock() + defer t.mu.Unlock() + + s, ok := t.selectors[selectorExpression] + if !ok { + return + } + s.etag = etag + + t.logger.Debug(fmt.Sprintf("recomputed ETag for selector '%s': %s", selectorExpression, etag)) +} diff --git a/flagd/pkg/service/flag-evaluation/ofrep/selector_version_tracker_test.go b/flagd/pkg/service/flag-evaluation/ofrep/selector_version_tracker_test.go new file mode 100644 index 000000000..2639a8f9d --- /dev/null +++ b/flagd/pkg/service/flag-evaluation/ofrep/selector_version_tracker_test.go @@ -0,0 +1,176 @@ +package ofrep + +import ( + "context" + "testing" + "time" + + "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" +) + +type mockFlagStore struct { + flags []model.Flag + watchCh chan struct{} +} + +func (m *mockFlagStore) GetAll(_ context.Context, _ *store.Selector) ([]model.Flag, model.Metadata, error) { + return m.flags, model.Metadata{}, nil +} + +func (m *mockFlagStore) WatchSelector(_ *store.Selector) <-chan struct{} { + return m.watchCh +} + +func TestNewSelectorVersionTracker(t *testing.T) { + log := logger.NewLogger(nil, false) + tracker := NewSelectorVersionTracker(log, nil, 0) + require.NotNil(t, tracker) + defer tracker.Close() +} + +func TestSelectorVersionTracker_Track(t *testing.T) { + log := logger.NewLogger(nil, false) + mockStore := &mockFlagStore{ + flags: []model.Flag{{Key: "test-flag", State: "ENABLED"}}, + watchCh: make(chan struct{}), + } + tracker := NewSelectorVersionTracker(log, mockStore, 0) + defer tracker.Close() + + selector := "flagSetId=test-set" + + // track the selector + etag := tracker.Track(selector) + require.NotEmpty(t, etag) + + // tracking again should return the same etag + etag2 := tracker.Track(selector) + assert.Equal(t, etag, etag2) + + // ETag should return the same value + assert.Equal(t, etag, tracker.ETag(selector)) + + // untracked selector should return empty etag + assert.Empty(t, tracker.ETag("flagSetId=different-set")) +} + +func TestSelectorVersionTracker_EmptySelector(t *testing.T) { + log := logger.NewLogger(nil, false) + mockStore := &mockFlagStore{ + flags: []model.Flag{{Key: "test-flag", State: "ENABLED"}}, + watchCh: make(chan struct{}), + } + tracker := NewSelectorVersionTracker(log, mockStore, 0) + defer tracker.Close() + + // test with empty selector + etag := tracker.Track("") + require.NotEmpty(t, etag) + + assert.Equal(t, etag, tracker.ETag("")) +} + +func TestSelectorVersionTracker_ContentBasedETags(t *testing.T) { + log := logger.NewLogger(nil, false) + + // two stores with the same content should produce the same ETag + flags := []model.Flag{{Key: "test-flag", State: "ENABLED", DefaultVariant: "on"}} + + store1 := &mockFlagStore{flags: flags, watchCh: make(chan struct{})} + store2 := &mockFlagStore{flags: flags, watchCh: make(chan struct{})} + + tracker1 := NewSelectorVersionTracker(log, store1, 0) + tracker2 := NewSelectorVersionTracker(log, store2, 0) + defer tracker1.Close() + defer tracker2.Close() + + etag1 := tracker1.Track("") + etag2 := tracker2.Track("") + + assert.Equal(t, etag1, etag2, "same content should produce same ETag across replicas") +} + +func TestSelectorVersionTracker_DifferentContentDifferentETags(t *testing.T) { + log := logger.NewLogger(nil, false) + + store1 := &mockFlagStore{ + flags: []model.Flag{{Key: "flag-a", State: "ENABLED"}}, + watchCh: make(chan struct{}), + } + store2 := &mockFlagStore{ + flags: []model.Flag{{Key: "flag-b", State: "ENABLED"}}, + watchCh: make(chan struct{}), + } + + tracker1 := NewSelectorVersionTracker(log, store1, 0) + tracker2 := NewSelectorVersionTracker(log, store2, 0) + defer tracker1.Close() + defer tracker2.Close() + + etag1 := tracker1.Track("") + etag2 := tracker2.Track("") + + // different content = different ETags + assert.NotEqual(t, etag1, etag2, "different content should produce different ETags") +} + +func TestSelectorVersionTracker_WatchBasedRecompute(t *testing.T) { + log := logger.NewLogger(nil, false) + + watchCh := make(chan struct{}) + mockStore := &mockFlagStore{ + flags: []model.Flag{{Key: "test-flag", State: "ENABLED"}}, + watchCh: watchCh, + } + + tracker := NewSelectorVersionTracker(log, mockStore, 0) + defer tracker.Close() + + // track a selector (this starts a watch goroutine) + selector := "flagSetId=test" + initialETag := tracker.Track(selector) + require.NotEmpty(t, initialETag) + + // change the flags + mockStore.flags = []model.Flag{{Key: "test-flag", State: "DISABLED"}} + + // simulate a store update + close(watchCh) + + time.Sleep(50 * time.Millisecond) + + // ETag should change because content changed + newETag := tracker.ETag(selector) + assert.NotEqual(t, initialETag, newETag, "ETag should change after content changes") +} + +func TestSelectorVersionTracker_Eviction(t *testing.T) { + log := logger.NewLogger(nil, false) + mockStore := &mockFlagStore{ + flags: []model.Flag{{Key: "test-flag", State: "ENABLED"}}, + watchCh: make(chan struct{}), + } + + // create tracker with capacity of 2 + tracker := NewSelectorVersionTracker(log, mockStore, 2) + defer tracker.Close() + + // track 2 selectors + tracker.Track("selector-1") + tracker.Track("selector-2") + + // both should be tracked + assert.NotEmpty(t, tracker.ETag("selector-1")) + assert.NotEmpty(t, tracker.ETag("selector-2")) + + // track a 3rd - should evict selector-1 (oldest) + tracker.Track("selector-3") + + assert.Empty(t, tracker.ETag("selector-1"), "selector-1 should be evicted") + assert.NotEmpty(t, tracker.ETag("selector-2")) + assert.NotEmpty(t, tracker.ETag("selector-3")) +}