Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
171 changes: 113 additions & 58 deletions src/databricks/sql/common/feature_flag.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,15 @@
import json
import math
import threading
import time
from ctypes import c_int32, c_int64
from dataclasses import dataclass, field
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Optional, List, Any, TYPE_CHECKING
from concurrent.futures import Future, ThreadPoolExecutor
from typing import Dict, Optional, List, Any

from databricks.sql.common.http import HttpMethod
from databricks.sql.common.url_utils import normalize_host_with_protocol

if TYPE_CHECKING:
from databricks.sql.client import Connection


@dataclass
class FeatureFlagEntry:
Expand Down Expand Up @@ -43,77 +42,126 @@ def from_dict(cls, data: Dict[str, Any]) -> "FeatureFlagsResponse":
REFRESH_BEFORE_EXPIRY_SECONDS = 10 # Start proactive refresh 10s before expiry


@dataclass
class _CacheState:
# Only values/coordination are shared; credentials and HTTP clients are not.
flags: Optional[Dict[str, str]] = None
ttl_seconds: int = DEFAULT_TTL_SECONDS
last_refresh_time: float = 0
lock: Any = field(default_factory=threading.RLock)
refresh: Optional[Future] = None


def _cache_key(host, headers):
workspace_id = (headers or {}).get("x-databricks-org-id")
return (
("workspace", workspace_id)
if workspace_id
else ("host", normalize_host_with_protocol(host).lower())
)


class FeatureFlagsContext:
"""
Manages fetching and caching of server-side feature flags for a connection.
Authenticated flag reader usable before any session/backend is opened.

1. The very first check for any flag is a synchronous, BLOCKING operation.
2. Subsequent refreshes (triggered near TTL expiry) are done asynchronously
in the background, returning stale data until the refresh completes.
"""

def __init__(
self, connection: "Connection", executor: ThreadPoolExecutor, http_client
self, host, executor, http_client, auth_provider, user_agent, headers, state
):
from databricks.sql import __version__

self._connection = connection
self._executor = executor # Used for ASYNCHRONOUS refreshes
self._lock = threading.RLock()

# Cache state: `None` indicates the cache has never been loaded.
self._flags: Optional[Dict[str, str]] = None
self._ttl_seconds: int = DEFAULT_TTL_SECONDS
self._last_refresh_time: float = 0
self._state = state
self._auth_provider = auth_provider
self._headers = {"User-Agent": user_agent, **headers}

endpoint_suffix = FEATURE_FLAGS_ENDPOINT_SUFFIX_FORMAT.format(__version__)
self._feature_flag_endpoint = (
normalize_host_with_protocol(self._connection.session.host)
+ endpoint_suffix
normalize_host_with_protocol(host) + endpoint_suffix
)

# Use the provided HTTP client
self._http_client = http_client

def _is_refresh_needed(self) -> bool:
"""Checks if the cache is due for a proactive background refresh."""
if self._flags is None:
if self._state.flags is None:
return False # Not eligible for refresh until loaded once.

refresh_threshold = self._last_refresh_time + (
self._ttl_seconds - REFRESH_BEFORE_EXPIRY_SECONDS
refresh_threshold = self._state.last_refresh_time + (
self._state.ttl_seconds - REFRESH_BEFORE_EXPIRY_SECONDS
)
return time.monotonic() > refresh_threshold

def get_flag_value(self, name: str, default_value: Any) -> Any:
def _get_value(self, name: str) -> Any:
"""
Checks if a feature is enabled.
Reads and parses a flag's JSON value.
- BLOCKS on the first call until flags are fetched.
- Returns cached values on subsequent calls, triggering non-blocking refreshes if needed.
"""
with self._lock:
with self._state.lock:
# If cache has never been loaded, perform a synchronous, blocking fetch.
if self._flags is None:
if self._state.flags is None:
self._refresh_flags()

# If a proactive background refresh is needed, start one. This is non-blocking.
elif self._is_refresh_needed():
# We don't check for an in-flight refresh; the executor queues the task, which is safe.
self._executor.submit(self._refresh_flags)

assert self._flags is not None
elif self._is_refresh_needed() and (
self._state.refresh is None or self._state.refresh.done()
):
self._state.refresh = self._executor.submit(self._refresh_flags)

# Now, return the value from the populated cache.
return self._flags.get(name, default_value)
raw = (self._state.flags or {}).get(name)
try:
return json.loads(raw) if raw is not None else None
except (TypeError, ValueError):
return None

def get_bool(self, name: str, default_value: bool = False) -> bool:
value = self._get_value(name)
return value if type(value) is bool else default_value

def _get_int(self, name, integer_type, default_value):
value = self._get_value(name)
if type(value) is int and integer_type(value).value == value:
return value
return default_value

def get_int32(self, name: str, default_value=None) -> Optional[int]:
return self._get_int(name, c_int32, default_value)

def get_int64(self, name: str, default_value=None) -> Optional[int]:
return self._get_int(name, c_int64, default_value)

def get_double(self, name: str, default_value=None) -> Optional[float]:
value = self._get_value(name)
if type(value) is int:
try:
value = float(value)
except OverflowError:
return default_value
return value if type(value) is float and math.isfinite(value) else default_value

def get_string(self, name: str, default_value=None) -> Optional[str]:
value = self._get_value(name)
return value if isinstance(value, str) else default_value

def get_string_list(self, name: str, default_value=None) -> Optional[List[str]]:
value = self._get_value(name)
if isinstance(value, list) and all(isinstance(item, str) for item in value):
return value
return default_value

def _refresh_flags(self):
"""Performs a synchronous network request to fetch and update flags."""
headers = {}
headers = dict(self._headers)
try:
# Authenticate the request
self._connection.session.auth_provider.add_headers(headers)
headers["User-Agent"] = self._connection.session.useragent_header
headers.update(self._connection.session.get_spog_headers())
self._auth_provider.add_headers(headers)

response = self._http_client.request(
HttpMethod.GET, self._feature_flag_endpoint, headers=headers, timeout=30
Expand All @@ -126,30 +174,30 @@ def _refresh_flags(self):
self._update_cache_from_response(ff_response)
else:
# On failure, initialize with an empty dictionary to prevent re-blocking.
if self._flags is None:
self._flags = {}
if self._state.flags is None:
self._state.flags = {}

except Exception as e:
except Exception:
# On exception, initialize with an empty dictionary to prevent re-blocking.
if self._flags is None:
self._flags = {}
if self._state.flags is None:
self._state.flags = {}

def _update_cache_from_response(self, ff_response: FeatureFlagsResponse):
"""Atomically updates the internal cache state from a successful server response."""
with self._lock:
self._flags = {flag.name: flag.value for flag in ff_response.flags}
with self._state.lock:
self._state.flags = {flag.name: flag.value for flag in ff_response.flags}
if ff_response.ttl_seconds is not None and ff_response.ttl_seconds > 0:
self._ttl_seconds = ff_response.ttl_seconds
self._last_refresh_time = time.monotonic()
self._state.ttl_seconds = ff_response.ttl_seconds
self._state.last_refresh_time = time.monotonic()


class FeatureFlagsContextFactory:
"""
Manages a singleton instance of FeatureFlagsContext per connection session.
Shares flag values per workspace, independent of telemetry/session lifetime.
Also manages a shared ThreadPoolExecutor for all background refresh operations.
"""

_context_map: Dict[str, FeatureFlagsContext] = {}
_context_map: Dict[tuple, _CacheState] = {}
_executor: Optional[ThreadPoolExecutor] = None
_lock = threading.Lock()

Expand All @@ -162,27 +210,34 @@ def _initialize(cls):
)

@classmethod
def get_instance(cls, connection: "Connection") -> FeatureFlagsContext:
"""Gets or creates a FeatureFlagsContext for the given connection."""
def get_instance(
cls, host, http_client, auth_provider, user_agent, headers=None
) -> FeatureFlagsContext:
"""Reuse the cache with this caller's authenticated transport, even pre-session."""
headers = {name.lower(): value for name, value in (headers or {}).items()}
with cls._lock:
cls._initialize()
assert cls._executor is not None

# Cache at HOST level - share feature flags across connections to same host
# Feature flags are per-host, not per-session
key = connection.session.host
key = _cache_key(host, headers)
if key not in cls._context_map:
cls._context_map[key] = FeatureFlagsContext(
connection, cls._executor, connection.session.http_client
)
return cls._context_map[key]
cls._context_map[key] = _CacheState()
return FeatureFlagsContext(
host,
cls._executor,
http_client,
auth_provider,
user_agent,
headers,
cls._context_map[key],
)

@classmethod
def remove_instance(cls, connection: "Connection"):
"""Removes the context for a given connection and shuts down the executor if no clients remain."""
def remove_instance(cls, host, headers=None):
"""Evicts a workspace's values and shuts down the executor if the cache is empty."""
with cls._lock:
# Use host as key to match get_instance
key = connection.session.host
headers = {name.lower(): value for name, value in (headers or {}).items()}
key = _cache_key(host, headers)
if key in cls._context_map:
cls._context_map.pop(key, None)

Expand Down
12 changes: 12 additions & 0 deletions src/databricks/sql/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,10 @@
from databricks.sql.backend.types import SessionId, BackendType
from databricks.sql.common.unified_http_client import UnifiedHttpClient
from databricks.sql.common.agent import detect as detect_agent
from databricks.sql.common.feature_flag import (
FeatureFlagsContext,
FeatureFlagsContextFactory,
)
from databricks.sql.telemetry.telemetry_client import TelemetryClientFactory

if TYPE_CHECKING:
Expand Down Expand Up @@ -134,6 +138,7 @@ def __init__(
# provider when an ``access_token`` is present, and ``None``
# otherwise (OAuth M2M/U2M resolve purely from the raw kwargs
# the bridge reads). The Thrift / SEA backends are unchanged.
self.feature_flags: Optional[FeatureFlagsContext] = None
if kwargs.get("use_kernel", False):
access_token = kwargs.get("access_token")
self.auth_provider = (
Expand All @@ -143,6 +148,13 @@ def __init__(
self.auth_provider = get_python_sql_connector_auth_provider(
server_hostname, http_client=self.http_client, **kwargs
)
self.feature_flags = FeatureFlagsContextFactory.get_instance(
self.host,
self.http_client,
self.auth_provider,
self.useragent_header,
self.get_spog_headers(),
)

self.backend = self._create_backend(
server_hostname,
Expand Down
9 changes: 4 additions & 5 deletions src/databricks/sql/telemetry/telemetry_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,6 @@
import uuid
import locale
from databricks.sql.telemetry.utils import BaseTelemetryClient
from databricks.sql.common.feature_flag import FeatureFlagsContextFactory
from databricks.sql.common.unified_http_client import UnifiedHttpClient
from databricks.sql.common.http import HttpMethod
from databricks.sql.exc import RequestError
Expand Down Expand Up @@ -133,12 +132,12 @@ def is_telemetry_enabled(connection: "Connection") -> bool:
if not connection.enable_telemetry:
return False

# Only fetch feature flags when enable_telemetry=True and not forced
context = FeatureFlagsContextFactory.get_instance(connection)
flag_value = context.get_flag_value(
context = connection.session.feature_flags
if context is None:
return False
return context.get_bool(
TelemetryHelper.TELEMETRY_FEATURE_FLAG_NAME, default_value=False
)
return str(flag_value).lower() == "true"


class NoopTelemetryClient(BaseTelemetryClient):
Expand Down
12 changes: 12 additions & 0 deletions tests/unit/conftest.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
from unittest.mock import patch

import pytest


@pytest.fixture(autouse=True)
def session_feature_flags():
# Unit connections must not contact connector-service. Cache/lifecycle tests
# exercise the real reader with a mocked HTTP client.
with patch("databricks.sql.session.FeatureFlagsContextFactory") as factory:
factory.get_instance.return_value.get_bool.return_value = False
yield factory
Loading
Loading