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/3] 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/3] 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 From e75abf3378d43fba7dae740d233217cea9e84a6c Mon Sep 17 00:00:00 2001 From: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> Date: Wed, 7 Oct 2026 00:30:28 +0000 Subject: [PATCH 3/3] fix(feature-flags): decouple cache lifetime from telemetry Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> --- connection.go | 16 ++-- connector.go | 42 +++++++---- connector_feature_flags_test.go | 126 ++++++++++++++++++++++++++++++++ telemetry/driver_integration.go | 13 +--- 4 files changed, 166 insertions(+), 31 deletions(-) create mode 100644 connector_feature_flags_test.go diff --git a/connection.go b/connection.go index 96e928b0..19de7f76 100644 --- a/connection.go +++ b/connection.go @@ -20,6 +20,7 @@ import ( "github.com/databricks/databricks-sql-go/internal/config" "github.com/databricks/databricks-sql-go/internal/debuglog" dbsqlerrint "github.com/databricks/databricks-sql-go/internal/errors" + "github.com/databricks/databricks-sql-go/internal/featureflags" "github.com/databricks/databricks-sql-go/internal/retry" "github.com/databricks/databricks-sql-go/internal/rows" "github.com/databricks/databricks-sql-go/logger" @@ -27,10 +28,11 @@ import ( ) type conn struct { - id string - cfg *config.Config - backend backend.Backend - telemetry *telemetry.Interceptor // Optional telemetry interceptor + id string + cfg *config.Config + backend backend.Backend + featureFlags *featureflags.Request + telemetry *telemetry.Interceptor // Optional telemetry interceptor } // tagStatementClosed returns a telemetry-only copy of a close-RPC error tagged @@ -75,7 +77,11 @@ 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, extractSpogHeaders(c.cfg.HTTPPath)["x-databricks-org-id"]) + telemetry.ReleaseForConnection(c.cfg.Host) + } + if c.featureFlags != nil { + featureflags.GetCache().Release(c.featureFlags.Host, c.featureFlags.WorkspaceID) + c.featureFlags = nil } if err != nil { diff --git a/connector.go b/connector.go index d8e030d2..2063dc73 100644 --- a/connector.go +++ b/connector.go @@ -23,6 +23,7 @@ import ( "github.com/databricks/databricks-sql-go/internal/client" "github.com/databricks/databricks-sql-go/internal/config" "github.com/databricks/databricks-sql-go/internal/debuglog" + "github.com/databricks/databricks-sql-go/internal/featureflags" "github.com/databricks/databricks-sql-go/internal/warehouse_cache" "github.com/databricks/databricks-sql-go/logger" "github.com/databricks/databricks-sql-go/telemetry" @@ -66,11 +67,19 @@ type federatedTokenAuthenticator struct { func (c *connector) Connect(ctx context.Context) (driver.Conn, error) { defer debuglog.Track(ctx, "connector.Connect", "host=%s", c.cfg.Host)() + flagRequest := c.featureFlagRequest() + if !c.cfg.UseKernel { + featureflags.GetCache().Acquire(flagRequest.Host, flagRequest.WorkspaceID) + } + // openSessionWithReydenFallback handles the session opening with automatic // recovery for Reyden / Real-Time warehouses that reject Thrift. It returns // the backend, latency, and error. be, sessionLatencyMs, err := c.openSessionWithReydenFallback(ctx) if err != nil { + if !c.cfg.UseKernel { + featureflags.GetCache().Release(flagRequest.Host, flagRequest.WorkspaceID) + } return nil, err } @@ -79,18 +88,10 @@ func (c *connector) Connect(ctx context.Context) (driver.Conn, error) { cfg: c.cfg, backend: be, } - log := logger.WithContext(conn.id, driverctx.CorrelationIdFromContext(ctx), "") - - // Extract SPOG routing headers from HTTPPath. When the workspace ID is - // available via ?o= or a cluster /o// path segment, - // wrap the HTTP client used for telemetry + feature-flag calls with a - // transport that injects x-databricks-org-id. Thrift routes via the URL so - // its own c.client doesn't need wrapping. - telemetryClient := c.client - spogHeaders := extractSpogHeaders(c.cfg.HTTPPath) - if len(spogHeaders) > 0 { - telemetryClient = withSpogHeaders(c.client, spogHeaders) + if !c.cfg.UseKernel { + conn.featureFlags = &flagRequest } + log := logger.WithContext(conn.id, driverctx.CorrelationIdFromContext(ctx), "") // Skip driver telemetry on the kernel path. The kernel owns query execution // below the driver backend, so keeping the Go telemetry interceptor active @@ -104,10 +105,10 @@ 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"], + WorkspaceID: flagRequest.WorkspaceID, DriverVersion: c.cfg.DriverVersion, - UserAgent: client.BuildUserAgent(c.cfg), - HTTPClient: telemetryClient, + UserAgent: flagRequest.UserAgent, + HTTPClient: flagRequest.HTTPClient, EnableTelemetry: c.cfg.EnableTelemetry, BatchSize: c.cfg.TelemetryBatchSize, FlushInterval: c.cfg.TelemetryFlushInterval, @@ -130,6 +131,19 @@ func (c *connector) Connect(ctx context.Context) (driver.Conn, error) { return conn, nil } +func (c *connector) featureFlagRequest() featureflags.Request { + headers := extractSpogHeaders(c.cfg.HTTPPath) + httpClient := c.client + if len(headers) > 0 { + httpClient = withSpogHeaders(httpClient, headers) + } + return featureflags.Request{ + Host: c.cfg.Host, WorkspaceID: headers["x-databricks-org-id"], + DriverVersion: c.cfg.DriverVersion, UserAgent: client.BuildUserAgent(c.cfg), + HTTPClient: httpClient, + } +} + // Driver returns underlying databricksDriver for compatibility with sql.DB Driver method func (c *connector) Driver() driver.Driver { return &databricksDriver{} diff --git a/connector_feature_flags_test.go b/connector_feature_flags_test.go new file mode 100644 index 00000000..5046b137 --- /dev/null +++ b/connector_feature_flags_test.go @@ -0,0 +1,126 @@ +package dbsql + +import ( + "context" + "errors" + "io" + "net/http" + "strings" + "sync/atomic" + "testing" + + "github.com/databricks/databricks-sql-go/internal/backend" + "github.com/databricks/databricks-sql-go/internal/backend/thrift" + "github.com/databricks/databricks-sql-go/internal/cli_service" + "github.com/databricks/databricks-sql-go/internal/client" + "github.com/databricks/databricks-sql-go/internal/config" + "github.com/databricks/databricks-sql-go/internal/featureflags" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type flagRoundTripper func(*http.Request) (*http.Response, error) + +func (f flagRoundTripper) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +func TestConnectionOwnsFlagsWithTelemetryDisabled(t *testing.T) { + var fetches atomic.Int32 + transport := flagRoundTripper(func(r *http.Request) (*http.Response, error) { + fetches.Add(1) + assert.Equal(t, "Bearer test-token", r.Header.Get("Authorization")) + assert.Equal(t, "123", r.Header.Get("X-Databricks-Org-Id")) + assert.Equal(t, "/api/2.0/connector-service/feature-flags/GOLANG/"+DriverVersion, r.URL.Path) + return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader( + `{"flags":[{"name":"sampleLimit","value":"42"}]}`)), Header: make(http.Header)}, nil + }) + driverConnector, err := NewConnector( + WithServerHostname("flags.example"), WithHTTPPath("/sql/1.0/warehouses/test?o=123"), + WithAccessToken("test-token"), WithTransport(transport), + func(cfg *config.Config) { cfg.EnableTelemetry = config.NewConfigValue(false) }, + ) + require.NoError(t, err) + c := driverConnector.(*connector) + ctx := context.Background() + flags := featureflags.GetCache() + request := c.featureFlagRequest() + c.thriftBackendFactory = func(ctx context.Context, cfg *config.Config, _ *http.Client) (backend.Backend, error) { + return thrift.NewForTest(&client.TestClient{ + FnOpenSession: func(context.Context, *cli_service.TOpenSessionReq) (*cli_service.TOpenSessionResp, error) { + return getTestSession(), nil + }, + FnCloseSession: func(context.Context, *cli_service.TCloseSessionReq) (*cli_service.TCloseSessionResp, error) { + return &cli_service.TCloseSessionResp{Status: &cli_service.TStatus{StatusCode: cli_service.TStatusCode_SUCCESS_STATUS}}, nil + }, + }, getTestSession(), cfg), nil + } + first, err := c.Connect(ctx) + require.NoError(t, err) + second, err := c.Connect(ctx) + require.NoError(t, err) + require.Nil(t, first.(*conn).telemetry) + require.Zero(t, fetches.Load(), "unused flags must not add connection requests") + value, err := flags.GetInt32(ctx, *first.(*conn).featureFlags, "sampleLimit") + require.NoError(t, err) + require.EqualValues(t, 42, value) + require.NoError(t, first.Close()) + value, err = flags.GetInt32(ctx, *second.(*conn).featureFlags, "sampleLimit") + require.NoError(t, err) + require.EqualValues(t, 42, value) + require.EqualValues(t, 1, fetches.Load()) + require.NoError(t, second.Close()) + flags.Acquire(request.Host, request.WorkspaceID) + defer flags.Release(request.Host, request.WorkspaceID) + _, err = flags.GetInt32(ctx, request, "sampleLimit") + require.NoError(t, err) + require.EqualValues(t, 2, fetches.Load(), "last close releases the workspace cache") +} + +func TestConnectionFlagFailureAndKernelPaths(t *testing.T) { + for _, scenario := range []string{"flag fetch fails", "backend fails", "explicit kernel"} { + t.Run(scenario, func(t *testing.T) { + var fetches atomic.Int32 + c, err := NewConnector(WithServerHostname("flags.example"), WithAccessToken("test-token"), + WithTransport(flagRoundTripper(func(*http.Request) (*http.Response, error) { + fetches.Add(1) + status := http.StatusOK + if scenario == "flag fetch fails" { + status = http.StatusServiceUnavailable + } + return &http.Response{StatusCode: status, Body: io.NopCloser(strings.NewReader(`{"flags":[]}`)), Header: make(http.Header)}, nil + })), WithUseKernel(scenario == "explicit kernel")) + require.NoError(t, err) + connector := c.(*connector) + connector.thriftBackendFactory = func(context.Context, *config.Config, *http.Client) (backend.Backend, error) { + if scenario == "backend fails" { + return nil, errors.New("backend failed") + } + return &fakeThriftBackend{sessionID: "test-session"}, nil + } + connector.kernelBackendFactory = func(context.Context, *config.Config) (backend.Backend, error) { + return &fakeKernelBackend{}, nil + } + connection, err := c.Connect(context.Background()) + if scenario == "backend fails" { + require.ErrorContains(t, err, "backend failed") + } else { + require.NoError(t, err) + require.Zero(t, fetches.Load()) + if scenario == "flag fetch fails" { + value, err := featureflags.GetCache().GetBool(context.Background(), *connection.(*conn).featureFlags, "missing") + require.Error(t, err) + require.False(t, value) + } + require.NoError(t, connection.Close()) + } + // No consumer is left after a failed open or close; a getter cannot fetch. + value, err := featureflags.GetCache().GetBool(context.Background(), connector.featureFlagRequest(), "missing") + require.NoError(t, err) + require.False(t, value) + if scenario == "flag fetch fails" { + require.EqualValues(t, 1, fetches.Load()) + } else { + require.Zero(t, fetches.Load()) + } + }) + } +} diff --git a/telemetry/driver_integration.go b/telemetry/driver_integration.go index 239f3616..110559c3 100644 --- a/telemetry/driver_integration.go +++ b/telemetry/driver_integration.go @@ -61,17 +61,12 @@ func InitializeForConnection(ctx context.Context, opts TelemetryInitOptions) *In cfg.FlushInterval = opts.FlushInterval } - // Get feature flag cache context FIRST (for reference counting) - flagCache := featureflags.GetCache() - flagCache.Acquire(opts.Host, opts.WorkspaceID) - // Check if telemetry should be enabled enabled := isTelemetryEnabled(ctx, cfg, featureflags.Request{ Host: opts.Host, WorkspaceID: opts.WorkspaceID, DriverVersion: opts.DriverVersion, UserAgent: opts.UserAgent, HTTPClient: opts.HTTPClient, }) if !enabled { - flagCache.Release(opts.Host, opts.WorkspaceID) return nil } @@ -79,8 +74,6 @@ func InitializeForConnection(ctx context.Context, opts TelemetryInitOptions) *In clientMgr := getClientManager() 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.Release(opts.Host, opts.WorkspaceID) return nil } @@ -92,12 +85,8 @@ func InitializeForConnection(ctx context.Context, opts TelemetryInitOptions) *In // // Parameters: // - host: Databricks host -func ReleaseForConnection(host string, workspaceID ...string) { +func ReleaseForConnection(host string) { // Release client manager reference clientMgr := getClientManager() _ = clientMgr.releaseClient(host) - - // Release feature flag cache reference - flagCache := featureflags.GetCache() - flagCache.Release(host, workspaceID...) }