Skip to content
195 changes: 135 additions & 60 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, Type, Union

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,141 @@ 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.
Comment thread
cathleeny marked this conversation as resolved.
- 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)
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)
Comment thread
cathleeny marked this conversation as resolved.

assert self._flags is not None

# 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 {"true": True, "false": False}.get(str(raw).lower())

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: str,
integer_type: Type[Union[c_int32, c_int64]],
default_value: Optional[int],
) -> Optional[int]:
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: Optional[int] = None
) -> Optional[int]:
return self._get_int(name, c_int32, default_value)

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

def get_double(
self, name: str, default_value: Optional[float] = 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: Optional[str] = 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: Optional[List[str]] = 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 +189,32 @@ 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.
Also manages a shared ThreadPoolExecutor for all background refresh operations.
Process-wide flag values per workspace and a shared refresh executor.

Both are created lazily and retained until process exit. Session close does
not evict values or shut down the executor, which other readers may still use.
"""

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

Expand All @@ -162,31 +227,41 @@ 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):
Comment thread
cathleeny marked this conversation as resolved.
"""Explicitly evict workspace values and stop the executor if the cache is empty.

Used for test/reset cleanup, not individual session teardown.
"""
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)

# If this was the last active context, clean up the thread pool.
# If no cached workspaces remain, clean up the thread pool.
if not cls._context_map and cls._executor is not None:
cls._executor.shutdown(wait=False)
cls._executor = None
22 changes: 20 additions & 2 deletions src/databricks/sql/session.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import logging
import re
from functools import cached_property
from typing import Dict, Tuple, List, Optional, Any, Type, TYPE_CHECKING

from databricks.sql.types import SSLOptions
Expand All @@ -13,6 +14,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,7 +139,8 @@ 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.
if kwargs.get("use_kernel", False):
self.use_kernel = kwargs.get("use_kernel", False)
if self.use_kernel:
access_token = kwargs.get("access_token")
self.auth_provider = (
AccessTokenAuthProvider(access_token) if access_token else None
Expand All @@ -155,6 +161,19 @@ def __init__(

self.protocol_version = None

@cached_property
def feature_flags(self) -> Optional[FeatureFlagsContext]:
"""Attach this session's transport to the shared cache only when needed."""
if self.use_kernel:
return None
return FeatureFlagsContextFactory.get_instance(
self.host,
self.http_client,
self.auth_provider,
self.useragent_header,
self.get_spog_headers(),
)

def _create_backend(
self,
server_hostname: str,
Expand All @@ -166,7 +185,6 @@ def _create_backend(
) -> DatabricksClient:
"""Create and return the appropriate backend client."""
self.use_sea = kwargs.get("use_sea", False)
self.use_kernel = kwargs.get("use_kernel", False)

if self.use_kernel and self.use_sea:
raise ValueError(
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