Skip to content
Open
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
2 changes: 1 addition & 1 deletion KERNEL_REV
Original file line number Diff line number Diff line change
@@ -1 +1 @@
80f2aee7d884994d7b0af9a9ea6078872859a9cd
2c958bfba476a0b0f165c829988959a32bf12f3e
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.
- 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)

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 None

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 Low — get_bool is stricter than the telemetry gate it replaces. The old path did str(flag_value).lower() == "true", which tolerated "True", "TRUE", and a quoted-string value. get_bool now requires the raw flag value to be a bare JSON boolean (json.loads(raw) must yield type(value) is bool). A server value of True/TRUE, or a JSON-quoted "\"true\"", now parses to a non-bool and silently returns the False default — disabling telemetry where it previously was enabled.

The unit mocks store the bare lowercase form (str(enabled).lower()), so tests don't exercise this. Worth confirming the connector-service contract guarantees the enableTelemetryForPythonDriver flag is emitted as a bare lowercase JSON boolean before relying on the stricter parse.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@cathleeny Is this comment valid?


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