From 3e1f6ee7a5f9e41126d5789b4044ab187c7bd8a7 Mon Sep 17 00:00:00 2001 From: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> Date: Tue, 6 Oct 2026 21:29:38 +0000 Subject: [PATCH 1/2] Generalize driver feature flag cache and honor server TTL Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> --- .golangci.yml | 1 + connection.go | 2 +- connector.go | 4 +- go.mod | 2 +- internal/featureflags/cache.go | 292 ++++++++++++++++++ .../featureflags/cache_test.go | 227 ++++++++++---- telemetry/config.go | 8 +- telemetry/config_test.go | 40 +-- telemetry/driver_integration.go | 23 +- telemetry/featureflag.go | 233 -------------- 10 files changed, 505 insertions(+), 327 deletions(-) create mode 100644 internal/featureflags/cache.go rename telemetry/featureflag_test.go => internal/featureflags/cache_test.go (57%) delete mode 100644 telemetry/featureflag.go diff --git a/.golangci.yml b/.golangci.yml index e7e7c614..d060a9ec 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -38,6 +38,7 @@ linters: - github.com/rs/zerolog - github.com/stretchr/testify - golang.org/x/oauth2 + - golang.org/x/sync/singleflight gosec: severity: low confidence: low diff --git a/connection.go b/connection.go index 6cde09a8..96e928b0 100644 --- a/connection.go +++ b/connection.go @@ -75,7 +75,7 @@ func (c *conn) Close() error { } c.telemetry.RecordOperation(ctx, c.id, "", telemetry.OperationTypeDeleteSession, time.Since(closeStart).Milliseconds(), telErr) _ = c.telemetry.Close(ctx) - telemetry.ReleaseForConnection(c.cfg.Host) + telemetry.ReleaseForConnection(c.cfg.Host, extractSpogHeaders(c.cfg.HTTPPath)["x-databricks-org-id"]) } if err != nil { diff --git a/connector.go b/connector.go index 99ab3646..d8e030d2 100644 --- a/connector.go +++ b/connector.go @@ -87,7 +87,8 @@ func (c *connector) Connect(ctx context.Context) (driver.Conn, error) { // transport that injects x-databricks-org-id. Thrift routes via the URL so // its own c.client doesn't need wrapping. telemetryClient := c.client - if spogHeaders := extractSpogHeaders(c.cfg.HTTPPath); len(spogHeaders) > 0 { + spogHeaders := extractSpogHeaders(c.cfg.HTTPPath) + if len(spogHeaders) > 0 { telemetryClient = withSpogHeaders(c.client, spogHeaders) } @@ -103,6 +104,7 @@ func (c *connector) Connect(ctx context.Context) (driver.Conn, error) { if !skipTelemetry { conn.telemetry = telemetry.InitializeForConnection(ctx, telemetry.TelemetryInitOptions{ Host: c.cfg.Host, + WorkspaceID: spogHeaders["x-databricks-org-id"], DriverVersion: c.cfg.DriverVersion, UserAgent: client.BuildUserAgent(c.cfg), HTTPClient: telemetryClient, diff --git a/go.mod b/go.mod index fd3c9168..8ce07746 100644 --- a/go.mod +++ b/go.mod @@ -12,6 +12,7 @@ require ( github.com/pierrec/lz4/v4 v4.1.15 github.com/stretchr/testify v1.8.1 golang.org/x/oauth2 v0.27.0 + golang.org/x/sync v0.22.0 gotest.tools/gotestsum v1.8.2 ) @@ -38,7 +39,6 @@ require ( github.com/zeebo/xxh3 v1.0.2 // indirect golang.org/x/crypto v0.55.0 // indirect golang.org/x/mod v0.40.0 // indirect - golang.org/x/sync v0.22.0 // indirect golang.org/x/telemetry v0.0.0-20260811182544-a038080d80e5 // indirect golang.org/x/term v0.45.0 // indirect golang.org/x/tools v0.49.0 // indirect diff --git a/internal/featureflags/cache.go b/internal/featureflags/cache.go new file mode 100644 index 00000000..91408f8f --- /dev/null +++ b/internal/featureflags/cache.go @@ -0,0 +1,292 @@ +// Package featureflags shares connector-service flags without requiring telemetry or a session. +package featureflags + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "sync" + "time" + + "github.com/databricks/databricks-sql-go/internal/client" + "golang.org/x/sync/singleflight" +) + +const ( + // featureFlagCacheDuration is the fallback when the server omits a valid TTL. + featureFlagCacheDuration = 15 * time.Minute + // featureFlagHTTPTimeout is the default timeout for feature flag HTTP requests + featureFlagHTTPTimeout = 10 * time.Second + // featureFlagEndpointPath is the path for feature flag endpoint + featureFlagEndpointPath = "/api/2.0/connector-service/feature-flags/GOLANG/" +) + +// Request provides caller-owned authenticated transport, usable before opening a session. +// Cache entries retain only values, never credentials or HTTP clients. +type Request struct { + Host, WorkspaceID, DriverVersion, UserAgent string + HTTPClient *http.Client +} + +// Cache shares values per workspace (normalized host when workspace ID is unavailable). +// Acquire a reference before reading, and Release it when the consumer closes. +type Cache struct { + mu sync.RWMutex + contexts map[string]*featureFlagContext +} + +// featureFlagContext holds feature flag state and reference count for a workspace. +type featureFlagContext struct { + mu sync.RWMutex // protects flags, lastFetched, cacheDuration + flags map[string]string + lastFetched time.Time + refCount int // protected by Cache.mu + cacheDuration time.Duration + fetch singleflight.Group +} + +var ( + flagCacheOnce sync.Once + flagCacheInstance *Cache +) + +// GetCache returns the process-wide cache shared by driver consumers. +func GetCache() *Cache { + flagCacheOnce.Do(func() { + flagCacheInstance = &Cache{ + contexts: make(map[string]*featureFlagContext), + } + }) + return flagCacheInstance +} + +func hostURL(host string) string { + if !strings.Contains(host, "://") { + host = "https://" + host + } + return strings.TrimRight(host, "/") +} + +func cacheKey(host string, workspaceID ...string) string { + if len(workspaceID) > 0 && workspaceID[0] != "" { + return "workspace:" + workspaceID[0] + } + normalized := strings.ToLower(hostURL(host)) + if strings.HasPrefix(normalized, "https://") { + normalized = strings.TrimSuffix(normalized, ":443") + } + return "host:" + normalized +} + +// Acquire increments the reference count without making a request. +func (c *Cache) Acquire(host string, workspaceID ...string) *featureFlagContext { + host = cacheKey(host, workspaceID...) + c.mu.Lock() + defer c.mu.Unlock() + + ctx, exists := c.contexts[host] + if !exists { + ctx = &featureFlagContext{ + cacheDuration: featureFlagCacheDuration, + } + c.contexts[host] = ctx + } + ctx.refCount++ + return ctx +} + +// Release removes the entry when its last consumer closes. +func (c *Cache) Release(host string, workspaceID ...string) { + host = cacheKey(host, workspaceID...) + c.mu.Lock() + defer c.mu.Unlock() + + if ctx, exists := c.contexts[host]; exists { + ctx.refCount-- + if ctx.refCount <= 0 { + delete(c.contexts, host) + } + } +} + +func (c *Cache) getValue(ctx context.Context, request Request, name string) (string, error) { + c.mu.RLock() + flagCtx, exists := c.contexts[cacheKey(request.Host, request.WorkspaceID)] + c.mu.RUnlock() + + if !exists { + return "", nil + } + + flagCtx.mu.RLock() + if !flagCtx.isExpired() { + value := flagCtx.flags[name] + flagCtx.mu.RUnlock() + return value, nil + } + flagCtx.mu.RUnlock() + + result := flagCtx.fetch.DoChan("", func() (any, error) { + flagCtx.mu.RLock() + if !flagCtx.isExpired() { + flags := flagCtx.flags + flagCtx.mu.RUnlock() + return flags, nil + } + flagCtx.mu.RUnlock() + flags, ttl, err := fetchFeatureFlags(ctx, request) + flagCtx.mu.Lock() + defer flagCtx.mu.Unlock() + if err == nil { + flagCtx.flags, flagCtx.cacheDuration = flags, ttl + flagCtx.lastFetched = time.Now() + } + if flagCtx.flags != nil { + return flagCtx.flags, nil // Retain stale values on a refresh failure. + } + return nil, err + }) + select { + case <-ctx.Done(): + return "", ctx.Err() + case fetched := <-result: + if fetched.Err != nil { + return "", fetched.Err + } + return fetched.Val.(map[string]string)[name], nil + } +} + +// isExpired returns true if the cache has expired. +func (c *featureFlagContext) isExpired() bool { + return c.flags == nil || time.Since(c.lastFetched) >= c.cacheDuration +} + +// GetBool defaults to false for absent/null flags; malformed values return false and an error. +func (c *Cache) GetBool(ctx context.Context, request Request, name string) (bool, error) { + return read[bool](c, ctx, request, name) +} + +func (c *Cache) GetInt32(ctx context.Context, request Request, name string) (int32, error) { + return read[int32](c, ctx, request, name) +} + +func (c *Cache) GetInt64(ctx context.Context, request Request, name string) (int64, error) { + return read[int64](c, ctx, request, name) +} + +func (c *Cache) GetDouble(ctx context.Context, request Request, name string) (float64, error) { + return read[float64](c, ctx, request, name) +} + +func (c *Cache) GetString(ctx context.Context, request Request, name string) (string, error) { + return read[string](c, ctx, request, name) +} + +func (c *Cache) GetStringList(ctx context.Context, request Request, name string) ([]string, error) { + items, err := read[[]*string](c, ctx, request, name) + if err != nil { + return nil, err + } + values := make([]string, len(items)) + for i, item := range items { + if item == nil { + return nil, fmt.Errorf("feature flag %q is not a string list", name) + } + values[i] = *item + } + return values, nil +} + +// Consumers choose the expected SAFE type. Non-boolean getters return an error +// for missing/null/invalid values, so callers can select a suitable default. +func read[T any](c *Cache, ctx context.Context, request Request, name string) (T, error) { + var value T + raw, err := c.getValue(ctx, request, name) + if err != nil { + return value, err + } + if raw == "" || strings.TrimSpace(raw) == "null" { + if _, boolean := any(value).(bool); boolean { + return value, nil + } + return value, fmt.Errorf("feature flag %q is missing or null", name) + } + if err := json.Unmarshal([]byte(raw), &value); err != nil { + var zero T + return zero, err + } + return value, nil +} + +func fetchFeatureFlags(ctx context.Context, request Request) (map[string]string, time.Duration, error) { + if request.HTTPClient == nil { + return nil, 0, fmt.Errorf("feature flags require an authenticated HTTP client") + } + // Add timeout to context if it doesn't have a deadline + if _, hasDeadline := ctx.Deadline(); !hasDeadline { + var cancel context.CancelFunc + ctx, cancel = context.WithTimeout(ctx, featureFlagHTTPTimeout) + defer cancel() + } + + // Construct endpoint URL using connector-service endpoint like JDBC + endpoint := fmt.Sprintf("%s%s%s", hostURL(request.Host), featureFlagEndpointPath, request.DriverVersion) + + // Feature-flag GET shares the same rate-limit group as /telemetry-ext on + // the server side, so a 429/503 here should also fail fast rather than + // being retried 5× by retryablehttp. + ctx = client.WithSkipTransientRetries(ctx) + + req, err := http.NewRequestWithContext(ctx, "GET", endpoint, nil) + if err != nil { + return nil, 0, fmt.Errorf("failed to create feature flag request: %w", err) + } + if request.UserAgent != "" { + req.Header.Set("User-Agent", request.UserAgent) + } + if request.WorkspaceID != "" { + req.Header.Set("X-Databricks-Org-Id", request.WorkspaceID) + } + + resp, err := request.HTTPClient.Do(req) + if err != nil { + return nil, 0, fmt.Errorf("failed to fetch feature flag: %w", err) + } + defer resp.Body.Close() //nolint:errcheck + + if resp.StatusCode != http.StatusOK { + // Read and discard body to allow HTTP connection reuse + _, _ = io.Copy(io.Discard, resp.Body) + return nil, 0, fmt.Errorf("feature flag check failed: %d", resp.StatusCode) + } + + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, 0, fmt.Errorf("failed to read feature flag response: %w", err) + } + + var result struct { + Flags []struct { + Name string `json:"name"` + Value string `json:"value"` + } `json:"flags"` + TTLSeconds int `json:"ttl_seconds"` + } + if err := json.Unmarshal(body, &result); err != nil { + return nil, 0, fmt.Errorf("failed to decode feature flag response: %w", err) + } + + flags := make(map[string]string, len(result.Flags)) + for _, flag := range result.Flags { + flags[flag.Name] = flag.Value + } + ttl := featureFlagCacheDuration + if seconds := result.TTLSeconds; seconds > 0 && int64(seconds) <= int64((1<<63-1)/time.Second) { + ttl = time.Duration(seconds) * time.Second + } + return flags, ttl, nil +} diff --git a/telemetry/featureflag_test.go b/internal/featureflags/cache_test.go similarity index 57% rename from telemetry/featureflag_test.go rename to internal/featureflags/cache_test.go index 6c789410..4ac115f2 100644 --- a/telemetry/featureflag_test.go +++ b/internal/featureflags/cache_test.go @@ -1,21 +1,28 @@ -package telemetry +package featureflags import ( "context" + "encoding/json" "net/http" "net/http/httptest" + "strconv" "sync" + "sync/atomic" "testing" "time" + + "github.com/stretchr/testify/require" ) +const featureFlagName = "databricks.partnerplatform.clientConfigsFeatureFlags.enableTelemetryForGoDriver" + func TestGetFeatureFlagCache_Singleton(t *testing.T) { // Reset singleton for testing flagCacheInstance = nil flagCacheOnce = sync.Once{} - cache1 := getFeatureFlagCache() - cache2 := getFeatureFlagCache() + cache1 := GetCache() + cache2 := GetCache() if cache1 != cache2 { t.Error("Expected singleton instances to be the same") @@ -23,14 +30,14 @@ func TestGetFeatureFlagCache_Singleton(t *testing.T) { } func TestFeatureFlagCache_GetOrCreateContext(t *testing.T) { - cache := &featureFlagCache{ + cache := &Cache{ contexts: make(map[string]*featureFlagContext), } host := "test-host.databricks.com" // First call should create context and increment refCount to 1 - ctx1 := cache.getOrCreateContext(host) + ctx1 := cache.Acquire(host) if ctx1 == nil { t.Fatal("Expected context to be created") } @@ -39,7 +46,7 @@ func TestFeatureFlagCache_GetOrCreateContext(t *testing.T) { } // Second call should reuse context and increment refCount to 2 - ctx2 := cache.getOrCreateContext(host) + ctx2 := cache.Acquire(host) if ctx2 != ctx1 { t.Error("Expected to get the same context instance") } @@ -54,19 +61,19 @@ func TestFeatureFlagCache_GetOrCreateContext(t *testing.T) { } func TestFeatureFlagCache_ReleaseContext(t *testing.T) { - cache := &featureFlagCache{ + cache := &Cache{ contexts: make(map[string]*featureFlagContext), } host := "test-host.databricks.com" // Create context with refCount = 2 - cache.getOrCreateContext(host) - cache.getOrCreateContext(host) + cache.Acquire(host) + cache.Acquire(host) // First release should decrement to 1 - cache.releaseContext(host) - ctx, exists := cache.contexts[host] + cache.Release(host) + ctx, exists := cache.contexts[cacheKey(host)] if !exists { t.Fatal("Expected context to still exist") } @@ -75,31 +82,31 @@ func TestFeatureFlagCache_ReleaseContext(t *testing.T) { } // Second release should remove context - cache.releaseContext(host) - _, exists = cache.contexts[host] + cache.Release(host) + _, exists = cache.contexts[cacheKey(host)] if exists { t.Error("Expected context to be removed when refCount reaches 0") } // Release non-existent context should not panic - cache.releaseContext("non-existent-host") + cache.Release("non-existent-host") } -func TestFeatureFlagCache_IsTelemetryEnabled_Cached(t *testing.T) { - cache := &featureFlagCache{ +func TestFeatureFlagCache_GetBool_Cached(t *testing.T) { + cache := &Cache{ contexts: make(map[string]*featureFlagContext), } host := "test-host.databricks.com" - ctx := cache.getOrCreateContext(host) + ctx := cache.Acquire(host) // Set cached value enabled := true - ctx.enabled = &enabled + ctx.flags = map[string]string{featureFlagName: strconv.FormatBool(enabled)} ctx.lastFetched = time.Now() // Should return cached value without HTTP call - result, err := cache.isTelemetryEnabled(context.Background(), host, "test-version", "test-ua", nil) + result, err := cache.GetBool(context.Background(), Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: nil}, featureFlagName) if err != nil { t.Errorf("Expected no error, got %v", err) } @@ -108,7 +115,7 @@ func TestFeatureFlagCache_IsTelemetryEnabled_Cached(t *testing.T) { } } -func TestFeatureFlagCache_IsTelemetryEnabled_Expired(t *testing.T) { +func TestFeatureFlagCache_GetBool_Expired(t *testing.T) { // Create mock server callCount := 0 server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -119,21 +126,21 @@ func TestFeatureFlagCache_IsTelemetryEnabled_Expired(t *testing.T) { })) defer server.Close() - cache := &featureFlagCache{ + cache := &Cache{ contexts: make(map[string]*featureFlagContext), } host := server.URL // Use full URL for testing - ctx := cache.getOrCreateContext(host) + ctx := cache.Acquire(host) // Set expired cached value enabled := false - ctx.enabled = &enabled + ctx.flags = map[string]string{featureFlagName: strconv.FormatBool(enabled)} ctx.lastFetched = time.Now().Add(-20 * time.Minute) // Expired // Should fetch fresh value httpClient := &http.Client{} - result, err := cache.isTelemetryEnabled(context.Background(), host, "test-version", "test-ua", httpClient) + result, err := cache.GetBool(context.Background(), Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: httpClient}, featureFlagName) if err != nil { t.Errorf("Expected no error, got %v", err) } @@ -145,20 +152,20 @@ func TestFeatureFlagCache_IsTelemetryEnabled_Expired(t *testing.T) { } // Verify cache was updated - if *ctx.enabled != true { + if ctx.flags[featureFlagName] != "true" { t.Error("Expected cache to be updated with new value") } } -func TestFeatureFlagCache_IsTelemetryEnabled_NoContext(t *testing.T) { - cache := &featureFlagCache{ +func TestFeatureFlagCache_GetBool_NoContext(t *testing.T) { + cache := &Cache{ contexts: make(map[string]*featureFlagContext), } host := "non-existent-host.databricks.com" // Should return false for non-existent context - result, err := cache.isTelemetryEnabled(context.Background(), host, "test-version", "test-ua", nil) + result, err := cache.GetBool(context.Background(), Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: nil}, featureFlagName) if err != nil { t.Errorf("Expected no error, got %v", err) } @@ -167,28 +174,28 @@ func TestFeatureFlagCache_IsTelemetryEnabled_NoContext(t *testing.T) { } } -func TestFeatureFlagCache_IsTelemetryEnabled_ErrorFallback(t *testing.T) { +func TestFeatureFlagCache_GetBool_ErrorFallback(t *testing.T) { // Create mock server that returns error server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusInternalServerError) })) defer server.Close() - cache := &featureFlagCache{ + cache := &Cache{ contexts: make(map[string]*featureFlagContext), } host := server.URL // Use full URL for testing - ctx := cache.getOrCreateContext(host) + ctx := cache.Acquire(host) // Set cached value enabled := true - ctx.enabled = &enabled + ctx.flags = map[string]string{featureFlagName: strconv.FormatBool(enabled)} ctx.lastFetched = time.Now().Add(-20 * time.Minute) // Expired // Should return cached value on error httpClient := &http.Client{} - result, err := cache.isTelemetryEnabled(context.Background(), host, "test-version", "test-ua", httpClient) + result, err := cache.GetBool(context.Background(), Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: httpClient}, featureFlagName) if err != nil { t.Errorf("Expected no error (fallback to cache), got %v", err) } @@ -197,23 +204,23 @@ func TestFeatureFlagCache_IsTelemetryEnabled_ErrorFallback(t *testing.T) { } } -func TestFeatureFlagCache_IsTelemetryEnabled_ErrorNoCache(t *testing.T) { +func TestFeatureFlagCache_GetBool_ErrorNoCache(t *testing.T) { // Create mock server that returns error server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusInternalServerError) })) defer server.Close() - cache := &featureFlagCache{ + cache := &Cache{ contexts: make(map[string]*featureFlagContext), } host := server.URL // Use full URL for testing - cache.getOrCreateContext(host) + cache.Acquire(host) // No cached value, should return error httpClient := &http.Client{} - result, err := cache.isTelemetryEnabled(context.Background(), host, "test-version", "test-ua", httpClient) + result, err := cache.GetBool(context.Background(), Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: httpClient}, featureFlagName) if err == nil { t.Error("Expected error when no cache available and fetch fails") } @@ -223,7 +230,7 @@ func TestFeatureFlagCache_IsTelemetryEnabled_ErrorNoCache(t *testing.T) { } func TestFeatureFlagCache_ConcurrentAccess(t *testing.T) { - cache := &featureFlagCache{ + cache := &Cache{ contexts: make(map[string]*featureFlagContext), } @@ -233,17 +240,17 @@ func TestFeatureFlagCache_ConcurrentAccess(t *testing.T) { var wg sync.WaitGroup wg.Add(numGoroutines) - // Concurrent getOrCreateContext + // Concurrent Acquire for i := 0; i < numGoroutines; i++ { go func() { defer wg.Done() - cache.getOrCreateContext(host) + cache.Acquire(host) }() } wg.Wait() // Verify refCount - ctx, exists := cache.contexts[host] + ctx, exists := cache.contexts[cacheKey(host)] if !exists { t.Fatal("Expected context to exist") } @@ -251,18 +258,18 @@ func TestFeatureFlagCache_ConcurrentAccess(t *testing.T) { t.Errorf("Expected refCount to be %d, got %d", numGoroutines, ctx.refCount) } - // Concurrent releaseContext + // Concurrent Release wg.Add(numGoroutines) for i := 0; i < numGoroutines; i++ { go func() { defer wg.Done() - cache.releaseContext(host) + cache.Release(host) }() } wg.Wait() // Verify context is removed - _, exists = cache.contexts[host] + _, exists = cache.contexts[cacheKey(host)] if exists { t.Error("Expected context to be removed after all releases") } @@ -271,28 +278,28 @@ func TestFeatureFlagCache_ConcurrentAccess(t *testing.T) { func TestFeatureFlagContext_IsExpired(t *testing.T) { tests := []struct { name string - enabled *bool + flags map[string]string fetched time.Time duration time.Duration want bool }{ { name: "no cache", - enabled: nil, + flags: nil, fetched: time.Time{}, duration: 15 * time.Minute, want: true, }, { name: "fresh cache", - enabled: boolPtr(true), + flags: map[string]string{featureFlagName: "true"}, fetched: time.Now(), duration: 15 * time.Minute, want: false, }, { name: "expired cache", - enabled: boolPtr(true), + flags: map[string]string{featureFlagName: "true"}, fetched: time.Now().Add(-20 * time.Minute), duration: 15 * time.Minute, want: true, @@ -302,7 +309,7 @@ func TestFeatureFlagContext_IsExpired(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { ctx := &featureFlagContext{ - enabled: tt.enabled, + flags: tt.flags, lastFetched: tt.fetched, cacheDuration: tt.duration, } @@ -330,11 +337,11 @@ func TestFetchFeatureFlag_Success(t *testing.T) { host := server.URL // Use full URL for testing httpClient := &http.Client{} - enabled, err := fetchFeatureFlag(context.Background(), host, "test-version", "test-ua", httpClient) + flags, _, err := fetchFeatureFlags(context.Background(), Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: httpClient}) if err != nil { t.Errorf("Expected no error, got %v", err) } - if !enabled { + if flags[featureFlagName] != "true" { t.Error("Expected feature flag to be enabled") } } @@ -352,7 +359,7 @@ func TestFetchFeatureFlag_SetsUserAgent(t *testing.T) { })) defer server.Close() - _, err := fetchFeatureFlag(context.Background(), server.URL, "9.9.9", wantUA, &http.Client{}) + _, _, err := fetchFeatureFlags(context.Background(), Request{Host: server.URL, DriverVersion: "9.9.9", UserAgent: wantUA, HTTPClient: &http.Client{}}) if err != nil { t.Fatalf("unexpected error: %v", err) } @@ -372,11 +379,11 @@ func TestFetchFeatureFlag_Disabled(t *testing.T) { host := server.URL // Use full URL for testing httpClient := &http.Client{} - enabled, err := fetchFeatureFlag(context.Background(), host, "test-version", "test-ua", httpClient) + flags, _, err := fetchFeatureFlags(context.Background(), Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: httpClient}) if err != nil { t.Errorf("Expected no error, got %v", err) } - if enabled { + if flags[featureFlagName] == "true" { t.Error("Expected feature flag to be disabled") } } @@ -392,11 +399,11 @@ func TestFetchFeatureFlag_FlagNotPresent(t *testing.T) { host := server.URL // Use full URL for testing httpClient := &http.Client{} - enabled, err := fetchFeatureFlag(context.Background(), host, "test-version", "test-ua", httpClient) + flags, _, err := fetchFeatureFlags(context.Background(), Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: httpClient}) if err != nil { t.Errorf("Expected no error, got %v", err) } - if enabled { + if flags[featureFlagName] == "true" { t.Error("Expected feature flag to be false when not present") } } @@ -410,7 +417,7 @@ func TestFetchFeatureFlag_HTTPError(t *testing.T) { host := server.URL // Use full URL for testing httpClient := &http.Client{} - _, err := fetchFeatureFlag(context.Background(), host, "test-version", "test-ua", httpClient) + _, _, err := fetchFeatureFlags(context.Background(), Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: httpClient}) if err == nil { t.Error("Expected error for HTTP 500") } @@ -427,7 +434,7 @@ func TestFetchFeatureFlag_InvalidJSON(t *testing.T) { host := server.URL // Use full URL for testing httpClient := &http.Client{} - _, err := fetchFeatureFlag(context.Background(), host, "test-version", "test-ua", httpClient) + _, _, err := fetchFeatureFlags(context.Background(), Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: httpClient}) if err == nil { t.Error("Expected error for invalid JSON") } @@ -446,13 +453,109 @@ func TestFetchFeatureFlag_ContextCancellation(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) cancel() // Cancel immediately - _, err := fetchFeatureFlag(ctx, host, "test-version", "test-ua", httpClient) + _, _, err := fetchFeatureFlags(ctx, Request{Host: host, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: httpClient}) if err == nil { t.Error("Expected error for cancelled context") } } -// Helper function to create bool pointer -func boolPtr(b bool) *bool { - return &b +func TestTypedFlags(t *testing.T) { + cache := &Cache{contexts: make(map[string]*featureFlagContext)} + entry := cache.Acquire("test-host") + entry.flags = map[string]string{ + "bool": "true", "int32": "2147483647", "int64": "9223372036854775807", + "double": "1.25", "string": `"hello"`, "list": `["a","b"]`, + "overflow": "9223372036854775808", "null": "null", "badlist": `["a",null]`, + } + entry.lastFetched = time.Now() + ctx, request := context.Background(), Request{Host: "test-host"} + b, err := cache.GetBool(ctx, request, "bool") + require.NoError(t, err) + require.True(t, b) + i32, err := cache.GetInt32(ctx, request, "int32") + require.NoError(t, err) + require.Equal(t, int32(2147483647), i32) + i64, err := cache.GetInt64(ctx, request, "int64") + require.NoError(t, err) + require.Equal(t, int64(9223372036854775807), i64) + d, err := cache.GetDouble(ctx, request, "double") + require.NoError(t, err) + require.Equal(t, 1.25, d) + s, err := cache.GetString(ctx, request, "string") + require.NoError(t, err) + require.Equal(t, "hello", s) + list, err := cache.GetStringList(ctx, request, "list") + require.NoError(t, err) + require.Equal(t, []string{"a", "b"}, list) + for _, name := range []string{"overflow", "double", "bool", "null", "missing"} { + _, err := cache.GetInt64(ctx, request, name) + require.Error(t, err, name) + } + _, err = cache.GetInt32(ctx, request, "int64") + require.Error(t, err) + _, err = cache.GetStringList(ctx, request, "badlist") + require.Error(t, err) + for _, name := range []string{"null", "missing"} { + b, err := cache.GetBool(ctx, request, name) + require.NoError(t, err) + require.False(t, b) + } +} + +func TestWorkspaceFetchAndRefresh(t *testing.T) { + var calls atomic.Int32 + var fail atomic.Bool + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls.Add(1) + if fail.Load() { + w.WriteHeader(http.StatusServiceUnavailable) + return + } + require.Equal(t, "GET", r.Method) + require.Equal(t, featureFlagEndpointPath+"test-version", r.URL.Path) + require.Equal(t, "test-ua", r.UserAgent()) + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "flags": []map[string]string{{"name": "flag", "value": r.Header.Get("X-Databricks-Org-Id")}}, + "ttl_seconds": 60, + }) + })) + defer server.Close() + cache := &Cache{contexts: make(map[string]*featureFlagContext)} + entry := cache.Acquire(server.URL, "1") + require.Same(t, entry, cache.Acquire("alias-host", "1")) + require.NotSame(t, entry, cache.Acquire(server.URL, "2")) + require.Same(t, cache.Acquire("TEST-host"), cache.Acquire("https://test-host/")) + request := Request{Host: server.URL, WorkspaceID: "1", DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: server.Client()} + var wg sync.WaitGroup + for range 20 { + wg.Add(1) + go func() { + defer wg.Done() + value, err := cache.GetInt64(context.Background(), request, "flag") + if err != nil || value != 1 { + t.Errorf("cold concurrent read = %d, %v", value, err) + } + }() + } + wg.Wait() + require.Equal(t, int32(1), calls.Load()) + require.Equal(t, time.Minute, entry.cacheDuration) + entry.lastFetched = time.Now().Add(-61 * time.Second) + value, err := cache.GetInt64(context.Background(), request, "flag") + require.NoError(t, err) + require.Equal(t, int64(1), value) + require.Equal(t, int32(2), calls.Load()) + entry.lastFetched = time.Now().Add(-61 * time.Second) + fail.Store(true) + value, err = cache.GetInt64(context.Background(), request, "flag") + require.NoError(t, err) + require.Equal(t, int64(1), value) // Stale value survives failed refresh. + fail.Store(false) + request.WorkspaceID = "2" + value, err = cache.GetInt64(context.Background(), request, "flag") + require.NoError(t, err) + require.Equal(t, int64(2), value) + cache.Release(server.URL, "2") + require.NotContains(t, cache.contexts, cacheKey(server.URL, "2")) } diff --git a/telemetry/config.go b/telemetry/config.go index ceb0ac21..355a0057 100644 --- a/telemetry/config.go +++ b/telemetry/config.go @@ -2,9 +2,10 @@ package telemetry import ( "context" - "net/http" "strconv" "time" + + "github.com/databricks/databricks-sql-go/internal/featureflags" ) // Config holds telemetry configuration. @@ -90,12 +91,13 @@ func ParseTelemetryConfig(params map[string]string) *Config { // (databricks.partnerplatform.clientConfigsFeatureFlags.enableTelemetryForGoDriver). // // In all other cases — explicit opt-out or server flag absent/unreachable — returns false. -func isTelemetryEnabled(ctx context.Context, cfg *Config, host string, driverVersion string, userAgent string, httpClient *http.Client) bool { +func isTelemetryEnabled(ctx context.Context, cfg *Config, request featureflags.Request) bool { if cfg.EnableTelemetry != nil { return *cfg.EnableTelemetry } - serverEnabled, err := getFeatureFlagCache().isTelemetryEnabled(ctx, host, driverVersion, userAgent, httpClient) + serverEnabled, err := featureflags.GetCache().GetBool(ctx, request, + "databricks.partnerplatform.clientConfigsFeatureFlags.enableTelemetryForGoDriver") if err != nil { return false } diff --git a/telemetry/config_test.go b/telemetry/config_test.go index 099d90a6..a51b99d5 100644 --- a/telemetry/config_test.go +++ b/telemetry/config_test.go @@ -6,6 +6,8 @@ import ( "net/http/httptest" "testing" "time" + + "github.com/databricks/databricks-sql-go/internal/featureflags" ) func TestDefaultConfig(t *testing.T) { @@ -138,6 +140,8 @@ func TestParseTelemetryConfig_AllParams(t *testing.T) { } } +func boolPtr(b bool) *bool { return &b } + // TestIsTelemetryEnabled_ExplicitOptOut: client sets enableTelemetry=false → // disabled even when server flag is true. Server is not consulted. func TestIsTelemetryEnabled_ExplicitOptOut(t *testing.T) { @@ -147,7 +151,7 @@ func TestIsTelemetryEnabled_ExplicitOptOut(t *testing.T) { })) defer server.Close() - result := isTelemetryEnabled(context.Background(), &Config{EnableTelemetry: boolPtr(false)}, server.URL, "test-version", "test-ua", &http.Client{Timeout: 5 * time.Second}) + result := isTelemetryEnabled(context.Background(), &Config{EnableTelemetry: boolPtr(false)}, featureflags.Request{Host: server.URL, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: &http.Client{Timeout: 5 * time.Second}}) if result { t.Error("Expected telemetry to be disabled when client sets enableTelemetry=false, got enabled") @@ -157,7 +161,7 @@ func TestIsTelemetryEnabled_ExplicitOptOut(t *testing.T) { // TestIsTelemetryEnabled_ExplicitOptIn: client sets enableTelemetry=true → // enabled without any server call (unreachable host proves no network call is made). func TestIsTelemetryEnabled_ExplicitOptIn(t *testing.T) { - result := isTelemetryEnabled(context.Background(), &Config{EnableTelemetry: boolPtr(true)}, "http://unreachable-host", "test-version", "test-ua", &http.Client{Timeout: 5 * time.Second}) + result := isTelemetryEnabled(context.Background(), &Config{EnableTelemetry: boolPtr(true)}, featureflags.Request{Host: "http://unreachable-host", DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: &http.Client{Timeout: 5 * time.Second}}) if !result { t.Error("Expected telemetry to be enabled when client sets enableTelemetry=true, got disabled") @@ -172,11 +176,11 @@ func TestIsTelemetryEnabled_ServerEnabled(t *testing.T) { })) defer server.Close() - flagCache := getFeatureFlagCache() - flagCache.getOrCreateContext(server.URL) - defer flagCache.releaseContext(server.URL) + flagCache := featureflags.GetCache() + flagCache.Acquire(server.URL) + defer flagCache.Release(server.URL) - result := isTelemetryEnabled(context.Background(), &Config{}, server.URL, "test-version", "test-ua", &http.Client{Timeout: 5 * time.Second}) + result := isTelemetryEnabled(context.Background(), &Config{}, featureflags.Request{Host: server.URL, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: &http.Client{Timeout: 5 * time.Second}}) if !result { t.Error("Expected telemetry to be enabled when server flag is true and EnableTelemetry is nil, got disabled") @@ -191,11 +195,11 @@ func TestIsTelemetryEnabled_ServerDisabled(t *testing.T) { })) defer server.Close() - flagCache := getFeatureFlagCache() - flagCache.getOrCreateContext(server.URL) - defer flagCache.releaseContext(server.URL) + flagCache := featureflags.GetCache() + flagCache.Acquire(server.URL) + defer flagCache.Release(server.URL) - result := isTelemetryEnabled(context.Background(), &Config{}, server.URL, "test-version", "test-ua", &http.Client{Timeout: 5 * time.Second}) + result := isTelemetryEnabled(context.Background(), &Config{}, featureflags.Request{Host: server.URL, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: &http.Client{Timeout: 5 * time.Second}}) if result { t.Error("Expected telemetry to be disabled when server flag is false and EnableTelemetry is nil, got enabled") @@ -209,11 +213,11 @@ func TestIsTelemetryEnabled_ServerError(t *testing.T) { })) defer server.Close() - flagCache := getFeatureFlagCache() - flagCache.getOrCreateContext(server.URL) - defer flagCache.releaseContext(server.URL) + flagCache := featureflags.GetCache() + flagCache.Acquire(server.URL) + defer flagCache.Release(server.URL) - result := isTelemetryEnabled(context.Background(), &Config{}, server.URL, "test-version", "test-ua", &http.Client{Timeout: 5 * time.Second}) + result := isTelemetryEnabled(context.Background(), &Config{}, featureflags.Request{Host: server.URL, DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: &http.Client{Timeout: 5 * time.Second}}) if result { t.Error("Expected telemetry to be disabled when server errors and EnableTelemetry is nil, got enabled") @@ -222,11 +226,11 @@ func TestIsTelemetryEnabled_ServerError(t *testing.T) { // TestIsTelemetryEnabled_ServerUnreachable: no DSN override, server unreachable → disabled. func TestIsTelemetryEnabled_ServerUnreachable(t *testing.T) { - flagCache := getFeatureFlagCache() - flagCache.getOrCreateContext("http://localhost:9999") - defer flagCache.releaseContext("http://localhost:9999") + flagCache := featureflags.GetCache() + flagCache.Acquire("http://localhost:9999") + defer flagCache.Release("http://localhost:9999") - result := isTelemetryEnabled(context.Background(), &Config{}, "http://localhost:9999", "test-version", "test-ua", &http.Client{Timeout: 1 * time.Second}) + result := isTelemetryEnabled(context.Background(), &Config{}, featureflags.Request{Host: "http://localhost:9999", DriverVersion: "test-version", UserAgent: "test-ua", HTTPClient: &http.Client{Timeout: 1 * time.Second}}) if result { t.Error("Expected telemetry to be disabled when server is unreachable and EnableTelemetry is nil, got enabled") diff --git a/telemetry/driver_integration.go b/telemetry/driver_integration.go index e33c7537..239f3616 100644 --- a/telemetry/driver_integration.go +++ b/telemetry/driver_integration.go @@ -6,6 +6,7 @@ import ( "time" "github.com/databricks/databricks-sql-go/internal/config" + "github.com/databricks/databricks-sql-go/internal/featureflags" ) // TelemetryInitOptions bundles the parameters for InitializeForConnection. @@ -13,6 +14,9 @@ type TelemetryInitOptions struct { // Host is the Databricks host. Host string + // WorkspaceID partitions flags when several workspaces share a SPOG host. + WorkspaceID string + // DriverVersion is the driver version string. DriverVersion string @@ -58,13 +62,16 @@ func InitializeForConnection(ctx context.Context, opts TelemetryInitOptions) *In } // Get feature flag cache context FIRST (for reference counting) - flagCache := getFeatureFlagCache() - flagCache.getOrCreateContext(opts.Host) + flagCache := featureflags.GetCache() + flagCache.Acquire(opts.Host, opts.WorkspaceID) // Check if telemetry should be enabled - enabled := isTelemetryEnabled(ctx, cfg, opts.Host, opts.DriverVersion, opts.UserAgent, opts.HTTPClient) + enabled := isTelemetryEnabled(ctx, cfg, featureflags.Request{ + Host: opts.Host, WorkspaceID: opts.WorkspaceID, DriverVersion: opts.DriverVersion, + UserAgent: opts.UserAgent, HTTPClient: opts.HTTPClient, + }) if !enabled { - flagCache.releaseContext(opts.Host) + flagCache.Release(opts.Host, opts.WorkspaceID) return nil } @@ -73,7 +80,7 @@ func InitializeForConnection(ctx context.Context, opts TelemetryInitOptions) *In telemetryClient := clientMgr.getOrCreateClient(opts.Host, opts.DriverVersion, opts.UserAgent, opts.HTTPClient, cfg) if telemetryClient == nil { // Client failed to start; release the flag cache ref we incremented above - flagCache.releaseContext(opts.Host) + flagCache.Release(opts.Host, opts.WorkspaceID) return nil } @@ -85,12 +92,12 @@ func InitializeForConnection(ctx context.Context, opts TelemetryInitOptions) *In // // Parameters: // - host: Databricks host -func ReleaseForConnection(host string) { +func ReleaseForConnection(host string, workspaceID ...string) { // Release client manager reference clientMgr := getClientManager() _ = clientMgr.releaseClient(host) // Release feature flag cache reference - flagCache := getFeatureFlagCache() - flagCache.releaseContext(host) + flagCache := featureflags.GetCache() + flagCache.Release(host, workspaceID...) } diff --git a/telemetry/featureflag.go b/telemetry/featureflag.go deleted file mode 100644 index 66c96439..00000000 --- a/telemetry/featureflag.go +++ /dev/null @@ -1,233 +0,0 @@ -package telemetry - -import ( - "context" - "encoding/json" - "fmt" - "io" - "net/http" - "sync" - "time" - - "github.com/databricks/databricks-sql-go/internal/client" -) - -const ( - // featureFlagCacheDuration is how long to cache feature flag values - featureFlagCacheDuration = 15 * time.Minute - // featureFlagHTTPTimeout is the default timeout for feature flag HTTP requests - featureFlagHTTPTimeout = 10 * time.Second - // featureFlagEndpointPath is the path for feature flag endpoint - featureFlagEndpointPath = "/api/2.0/connector-service/feature-flags/GOLANG/" - // featureFlagName is the name of the Go driver telemetry feature flag - featureFlagName = "databricks.partnerplatform.clientConfigsFeatureFlags.enableTelemetryForGoDriver" -) - -// featureFlagCache manages feature flag state per host with reference counting. -// This prevents rate limiting by caching feature flag responses. -type featureFlagCache struct { - mu sync.RWMutex - contexts map[string]*featureFlagContext -} - -// featureFlagContext holds feature flag state and reference count for a host. -type featureFlagContext struct { - mu sync.RWMutex // protects enabled, lastFetched, fetching - enabled *bool - lastFetched time.Time - refCount int // protected by featureFlagCache.mu - cacheDuration time.Duration - fetching bool // true if a fetch is in progress -} - -var ( - flagCacheOnce sync.Once - flagCacheInstance *featureFlagCache -) - -// getFeatureFlagCache returns the singleton instance. -func getFeatureFlagCache() *featureFlagCache { - flagCacheOnce.Do(func() { - flagCacheInstance = &featureFlagCache{ - contexts: make(map[string]*featureFlagContext), - } - }) - return flagCacheInstance -} - -// getOrCreateContext gets or creates a feature flag context for the host. -// Increments reference count. -func (c *featureFlagCache) getOrCreateContext(host string) *featureFlagContext { - c.mu.Lock() - defer c.mu.Unlock() - - ctx, exists := c.contexts[host] - if !exists { - ctx = &featureFlagContext{ - cacheDuration: featureFlagCacheDuration, - } - c.contexts[host] = ctx - } - ctx.refCount++ - return ctx -} - -// releaseContext decrements reference count for the host. -// Removes context when ref count reaches zero. -func (c *featureFlagCache) releaseContext(host string) { - c.mu.Lock() - defer c.mu.Unlock() - - if ctx, exists := c.contexts[host]; exists { - ctx.refCount-- - if ctx.refCount <= 0 { - delete(c.contexts, host) - } - } -} - -// isTelemetryEnabled checks if telemetry is enabled for the host. -// Uses cached value if available and not expired. -func (c *featureFlagCache) isTelemetryEnabled(ctx context.Context, host string, driverVersion string, userAgent string, httpClient *http.Client) (bool, error) { - c.mu.RLock() - flagCtx, exists := c.contexts[host] - c.mu.RUnlock() - - if !exists { - return false, nil - } - - // Fast path: check cache under read lock. - flagCtx.mu.RLock() - if flagCtx.enabled != nil && time.Since(flagCtx.lastFetched) < flagCtx.cacheDuration { - enabled := *flagCtx.enabled - flagCtx.mu.RUnlock() - return enabled, nil - } - if flagCtx.fetching { - if flagCtx.enabled != nil { - enabled := *flagCtx.enabled - flagCtx.mu.RUnlock() - return enabled, nil - } - flagCtx.mu.RUnlock() - return false, nil - } - flagCtx.mu.RUnlock() - - // Slow path: need a write lock to set fetching=true. - // Re-check all conditions under write lock (double-checked locking) to avoid - // a data race and to prevent duplicate fetches from concurrent goroutines. - flagCtx.mu.Lock() - if flagCtx.enabled != nil && time.Since(flagCtx.lastFetched) < flagCtx.cacheDuration { - enabled := *flagCtx.enabled - flagCtx.mu.Unlock() - return enabled, nil - } - if flagCtx.fetching { - if flagCtx.enabled != nil { - enabled := *flagCtx.enabled - flagCtx.mu.Unlock() - return enabled, nil - } - flagCtx.mu.Unlock() - return false, nil - } - flagCtx.fetching = true - flagCtx.mu.Unlock() - - // Fetch fresh value (outside lock so other readers are not blocked). - enabled, err := fetchFeatureFlag(ctx, host, driverVersion, userAgent, httpClient) - - // Update cache. - flagCtx.mu.Lock() - flagCtx.fetching = false - if err == nil { - flagCtx.enabled = &enabled - flagCtx.lastFetched = time.Now() - } - result := false - var returnErr error - if err != nil { - if flagCtx.enabled != nil { - result = *flagCtx.enabled // Return stale cached value on error - } else { - returnErr = err - } - } else { - result = enabled - } - flagCtx.mu.Unlock() - - return result, returnErr -} - -// isExpired returns true if the cache has expired. -func (c *featureFlagContext) isExpired() bool { - return c.enabled == nil || time.Since(c.lastFetched) > c.cacheDuration -} - -// fetchFeatureFlag fetches the feature flag value from Databricks. -func fetchFeatureFlag(ctx context.Context, host string, driverVersion string, userAgent string, httpClient *http.Client) (bool, error) { - // Add timeout to context if it doesn't have a deadline - if _, hasDeadline := ctx.Deadline(); !hasDeadline { - var cancel context.CancelFunc - ctx, cancel = context.WithTimeout(ctx, featureFlagHTTPTimeout) - defer cancel() - } - - // Construct endpoint URL using connector-service endpoint like JDBC - hostURL := ensureHTTPScheme(host) - endpoint := fmt.Sprintf("%s%s%s", hostURL, featureFlagEndpointPath, driverVersion) - - // Feature-flag GET shares the same rate-limit group as /telemetry-ext on - // the server side, so a 429/503 here should also fail fast rather than - // being retried 5× by retryablehttp. - ctx = client.WithSkipTransientRetries(ctx) - - req, err := http.NewRequestWithContext(ctx, "GET", endpoint, nil) - if err != nil { - return false, fmt.Errorf("failed to create feature flag request: %w", err) - } - if userAgent != "" { - req.Header.Set("User-Agent", userAgent) - } - - resp, err := httpClient.Do(req) - if err != nil { - return false, fmt.Errorf("failed to fetch feature flag: %w", err) - } - defer resp.Body.Close() //nolint:errcheck - - if resp.StatusCode != http.StatusOK { - // Read and discard body to allow HTTP connection reuse - _, _ = io.Copy(io.Discard, resp.Body) - return false, fmt.Errorf("feature flag check failed: %d", resp.StatusCode) - } - - body, err := io.ReadAll(resp.Body) - if err != nil { - return false, fmt.Errorf("failed to read feature flag response: %w", err) - } - - var result struct { - Flags []struct { - Name string `json:"name"` - Value string `json:"value"` - } `json:"flags"` - TTLSeconds int `json:"ttl_seconds"` - } - if err := json.Unmarshal(body, &result); err != nil { - return false, fmt.Errorf("failed to decode feature flag response: %w", err) - } - - // Look for Go driver telemetry feature flag - for _, flag := range result.Flags { - if flag.Name == featureFlagName { - enabled := flag.Value == "true" - return enabled, nil - } - } - - return false, nil -} From ef1968bcef1341a98dd143eec936625965d63f3f Mon Sep 17 00:00:00 2001 From: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> Date: Tue, 6 Oct 2026 22:55:45 +0000 Subject: [PATCH 2/2] fix(feature-flags): keep stale readers non-blocking during refresh Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> --- internal/featureflags/cache.go | 18 +++++++++-- internal/featureflags/cache_test.go | 50 +++++++++++++++++++++++++++++ 2 files changed, 65 insertions(+), 3 deletions(-) diff --git a/internal/featureflags/cache.go b/internal/featureflags/cache.go index 91408f8f..07f72161 100644 --- a/internal/featureflags/cache.go +++ b/internal/featureflags/cache.go @@ -46,6 +46,7 @@ type featureFlagContext struct { refCount int // protected by Cache.mu cacheDuration time.Duration fetch singleflight.Group + refresh sync.Mutex // concurrent stale readers skip an in-flight refresh } var ( @@ -122,14 +123,15 @@ func (c *Cache) getValue(ctx context.Context, request Request, name string) (str } flagCtx.mu.RLock() + cachedFlags := flagCtx.flags if !flagCtx.isExpired() { - value := flagCtx.flags[name] + value := cachedFlags[name] flagCtx.mu.RUnlock() return value, nil } flagCtx.mu.RUnlock() - result := flagCtx.fetch.DoChan("", func() (any, error) { + load := func() (map[string]string, error) { flagCtx.mu.RLock() if !flagCtx.isExpired() { flags := flagCtx.flags @@ -148,7 +150,17 @@ func (c *Cache) getValue(ctx context.Context, request Request, name string) (str return flagCtx.flags, nil // Retain stale values on a refresh failure. } return nil, err - }) + } + if cachedFlags != nil { + if !flagCtx.refresh.TryLock() { + return cachedFlags[name], nil + } + defer flagCtx.refresh.Unlock() + flags, err := load() + return flags[name], err + } + + result := flagCtx.fetch.DoChan("", func() (any, error) { return load() }) select { case <-ctx.Done(): return "", ctx.Err() diff --git a/internal/featureflags/cache_test.go b/internal/featureflags/cache_test.go index 4ac115f2..f5d86492 100644 --- a/internal/featureflags/cache_test.go +++ b/internal/featureflags/cache_test.go @@ -502,6 +502,56 @@ func TestTypedFlags(t *testing.T) { } } +func TestStaleReadersDoNotWaitForRefresh(t *testing.T) { + started := make(chan struct{}) + allowRefresh := make(chan struct{}) + unblockRefresh := sync.OnceFunc(func() { close(allowRefresh) }) + var calls atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if calls.Add(1) == 1 { + close(started) + } + <-allowRefresh + _ = json.NewEncoder(w).Encode(map[string]any{ + "flags": []map[string]string{{"name": "flag", "value": "true"}}, + "ttl_seconds": 60, + }) + })) + defer server.Close() + defer unblockRefresh() + + cache := &Cache{contexts: make(map[string]*featureFlagContext)} + entry := cache.Acquire(server.URL) + entry.flags = map[string]string{"flag": "false"} + entry.lastFetched = time.Now().Add(-time.Hour) + request := Request{Host: server.URL, DriverVersion: "test-version", HTTPClient: server.Client()} + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + refreshed := make(chan error, 1) + go func() { + _, err := cache.GetBool(ctx, request, "flag") + refreshed <- err + }() + select { + case <-started: + case <-ctx.Done(): + t.Fatal("refresh did not start") + } + + readerCtx, cancelReader := context.WithTimeout(ctx, 100*time.Millisecond) + defer cancelReader() + value, err := cache.GetBool(readerCtx, request, "flag") + require.NoError(t, err, "a stale reader must not wait for the refresh") + require.False(t, value) + + unblockRefresh() + require.NoError(t, <-refreshed) + value, err = cache.GetBool(ctx, request, "flag") + require.NoError(t, err) + require.True(t, value) + require.Equal(t, int32(1), calls.Load()) +} + func TestWorkspaceFetchAndRefresh(t *testing.T) { var calls atomic.Int32 var fail atomic.Bool