From 1ce7f5b01731e6adcd7e21d0055c443a5416300a Mon Sep 17 00:00:00 2001 From: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> Date: Tue, 6 Oct 2026 21:29:36 +0000 Subject: [PATCH 1/6] Generalize driver feature flag cache and typed reads Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> --- src/databricks/sql/common/feature_flag.py | 168 ++++++++++++------ .../sql/telemetry/telemetry_client.py | 12 +- tests/unit/test_telemetry.py | 102 ++++++++++- 3 files changed, 216 insertions(+), 66 deletions(-) diff --git a/src/databricks/sql/common/feature_flag.py b/src/databricks/sql/common/feature_flag.py index 0b2c7490b..70ab9e349 100644 --- a/src/databricks/sql/common/feature_flag.py +++ b/src/databricks/sql/common/feature_flag.py @@ -1,16 +1,14 @@ import json +import math import threading import time 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: @@ -43,9 +41,28 @@ 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 @@ -53,23 +70,18 @@ class FeatureFlagsContext: """ 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 @@ -77,43 +89,78 @@ def __init__( 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 + 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, bits, default_value): + value = self._get_value(name) + if type(value) is int and -(2 ** (bits - 1)) <= value < 2 ** (bits - 1): + return value + return default_value + + def get_int32(self, name: str, default_value=None) -> Optional[int]: + return self._get_int(name, 32, default_value) - # Now, return the value from the populated cache. - return self._flags.get(name, default_value) + def get_int64(self, name: str, default_value=None) -> Optional[int]: + return self._get_int(name, 64, default_value) + + def get_double(self, name: str, default_value=None) -> Optional[float]: + value = self._get_value(name) + try: + if type(value) in (int, float) and math.isfinite(value): + return float(value) + except OverflowError: + pass + return 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 @@ -126,30 +173,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() @@ -162,27 +209,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) diff --git a/src/databricks/sql/telemetry/telemetry_client.py b/src/databricks/sql/telemetry/telemetry_client.py index 2051fb2f8..6341066fe 100644 --- a/src/databricks/sql/telemetry/telemetry_client.py +++ b/src/databricks/sql/telemetry/telemetry_client.py @@ -134,11 +134,17 @@ def is_telemetry_enabled(connection: "Connection") -> bool: return False # Only fetch feature flags when enable_telemetry=True and not forced - context = FeatureFlagsContextFactory.get_instance(connection) - flag_value = context.get_flag_value( + session = connection.session + context = FeatureFlagsContextFactory.get_instance( + session.host, + session.http_client, + session.auth_provider, + session.useragent_header, + session.get_spog_headers(), + ) + return context.get_bool( TelemetryHelper.TELEMETRY_FEATURE_FLAG_NAME, default_value=False ) - return str(flag_value).lower() == "true" class NoopTelemetryClient(BaseTelemetryClient): diff --git a/tests/unit/test_telemetry.py b/tests/unit/test_telemetry.py index 5663fa119..1beaf51e4 100644 --- a/tests/unit/test_telemetry.py +++ b/tests/unit/test_telemetry.py @@ -1073,16 +1073,24 @@ def reset_factory(self): ) def test_host_level_caching(self, hosts, expected_contexts): """Test that contexts are cached by host correctly.""" + contexts = [] for host in hosts: conn = MagicMock() conn.session.host = host conn.session.http_client = MagicMock() - contexts.append(FeatureFlagsContextFactory.get_instance(conn)) + contexts.append( + FeatureFlagsContextFactory.get_instance( + host, + conn.session.http_client, + conn.session.auth_provider, + "test-agent", + ) + ) assert len(FeatureFlagsContextFactory._context_map) == expected_contexts if expected_contexts == 1: - assert all(ctx is contexts[0] for ctx in contexts) + assert all(ctx._state is contexts[0]._state for ctx in contexts) def test_remove_instance_and_executor_cleanup(self): """Test removal uses host key and cleans up executor when empty.""" @@ -1094,14 +1102,96 @@ def test_remove_instance_and_executor_cleanup(self): conn2.session.host = "host2.com" conn2.session.http_client = MagicMock() - FeatureFlagsContextFactory.get_instance(conn1) - FeatureFlagsContextFactory.get_instance(conn2) + FeatureFlagsContextFactory.get_instance("host1.com", conn1, conn1, "test-agent") + FeatureFlagsContextFactory.get_instance("host2.com", conn2, conn2, "test-agent") assert FeatureFlagsContextFactory._executor is not None - FeatureFlagsContextFactory.remove_instance(conn1) + FeatureFlagsContextFactory.remove_instance("host1.com") assert len(FeatureFlagsContextFactory._context_map) == 1 assert FeatureFlagsContextFactory._executor is not None - FeatureFlagsContextFactory.remove_instance(conn2) + FeatureFlagsContextFactory.remove_instance("host2.com") assert len(FeatureFlagsContextFactory._context_map) == 0 assert FeatureFlagsContextFactory._executor is None + + @pytest.mark.parametrize( + "getter,raw,expected", + [ + ("bool", "true", True), + ("bool", '"true"', False), + ("int32", "2147483647", 2147483647), + ("int32", "2147483648", None), + ("int64", "9223372036854775807", 9223372036854775807), + ("int64", "-9223372036854775808", -9223372036854775808), + ("int64", "9223372036854775808", None), + ("int64", "true", None), + ("double", "1.25", 1.25), + ("double", "1e400", None), + ("string", '"hello"', "hello"), + ("string", "null", None), + ("string_list", '["a", "b"]', ["a", "b"]), + ("string_list", '["a", null]', None), + ("string", "invalid", None), + ], + ) + def test_typed_reads_without_session(self, getter, raw, expected): + http = MagicMock() + http.request.return_value = MagicMock( + status=200, + data=json.dumps( + { + "flags": [{"name": "flag", "value": raw}], + "ttl_seconds": 60, + } + ).encode(), + ) + reader = FeatureFlagsContextFactory.get_instance( + "test-host", http, AccessTokenAuthProvider("token"), "agent" + ) + assert getattr(reader, "get_" + getter)("flag") == expected + assert reader.get_string("missing", "fallback") == "fallback" + http.request.assert_called_once() + assert ( + http.request.call_args.kwargs["headers"]["Authorization"] == "Bearer token" + ) + + def test_workspace_cache_and_refresh_use_current_caller(self): + from concurrent.futures import Future + from databricks.sql.common.feature_flag import ( + FeatureFlagsResponse, + FeatureFlagEntry, + ) + + def reader(host, workspace): + return FeatureFlagsContextFactory.get_instance( + host, + MagicMock(), + AccessTokenAuthProvider("token"), + "agent", + {"X-Databricks-Org-Id": workspace}, + ) + + first = reader("test-host", "1") + current = reader("alias-host", "1") + other = reader("test-host", "2") + assert current._state is first._state + assert other._state is not first._state + first._update_cache_from_response( + FeatureFlagsResponse([FeatureFlagEntry("flag", "true")], ttl_seconds=60) + ) + assert current.get_bool("flag") is True + current._http_client.request.assert_not_called() + pending = Future() + with patch.object(current._executor, "submit", return_value=pending) as submit: + with patch( + "databricks.sql.common.feature_flag.time.monotonic", + return_value=first._state.last_refresh_time + 61, + ): + assert current.get_bool("flag") is True + assert first.get_bool("flag") is True + submit.assert_called_once_with(current._refresh_flags) + current._http_client.request.side_effect = RuntimeError("unavailable") + current._refresh_flags() + assert current.get_bool("flag") is True + current._http_client.request.assert_called_once() + first._http_client.request.assert_not_called() From 0ea879acbb832c68e7e4090d044656ae8b6a201d Mon Sep 17 00:00:00 2001 From: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> Date: Tue, 6 Oct 2026 22:52:24 +0000 Subject: [PATCH 2/6] refactor(feature-flags): use standard integer types for validation Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> --- src/databricks/sql/common/feature_flag.py | 21 +++++++++++---------- 1 file changed, 11 insertions(+), 10 deletions(-) diff --git a/src/databricks/sql/common/feature_flag.py b/src/databricks/sql/common/feature_flag.py index 70ab9e349..fb1fad527 100644 --- a/src/databricks/sql/common/feature_flag.py +++ b/src/databricks/sql/common/feature_flag.py @@ -2,6 +2,7 @@ import math import threading import time +from ctypes import c_int32, c_int64 from dataclasses import dataclass, field from concurrent.futures import Future, ThreadPoolExecutor from typing import Dict, Optional, List, Any @@ -124,26 +125,26 @@ 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, bits, default_value): + def _get_int(self, name, integer_type, default_value): value = self._get_value(name) - if type(value) is int and -(2 ** (bits - 1)) <= value < 2 ** (bits - 1): + 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, 32, default_value) + 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, 64, default_value) + 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) - try: - if type(value) in (int, float) and math.isfinite(value): - return float(value) - except OverflowError: - pass - return default_value + 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) From 9fda19d20e86d217971b1555a301cf1f4700e7e1 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/6] fix(feature-flags): decouple cache setup from telemetry Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> --- src/databricks/sql/session.py | 12 +++ .../sql/telemetry/telemetry_client.py | 13 +-- tests/unit/conftest.py | 12 +++ tests/unit/test_feature_flag_lifecycle.py | 97 +++++++++++++++++++ tests/unit/test_session.py | 12 ++- tests/unit/test_telemetry.py | 8 ++ 6 files changed, 142 insertions(+), 12 deletions(-) create mode 100644 tests/unit/conftest.py create mode 100644 tests/unit/test_feature_flag_lifecycle.py diff --git a/src/databricks/sql/session.py b/src/databricks/sql/session.py index db512dd0b..eba208c61 100644 --- a/src/databricks/sql/session.py +++ b/src/databricks/sql/session.py @@ -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: @@ -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 = ( @@ -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, diff --git a/src/databricks/sql/telemetry/telemetry_client.py b/src/databricks/sql/telemetry/telemetry_client.py index 6341066fe..fed2001cb 100644 --- a/src/databricks/sql/telemetry/telemetry_client.py +++ b/src/databricks/sql/telemetry/telemetry_client.py @@ -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 @@ -133,15 +132,9 @@ 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 - session = connection.session - context = FeatureFlagsContextFactory.get_instance( - session.host, - session.http_client, - session.auth_provider, - session.useragent_header, - session.get_spog_headers(), - ) + context = connection.session.feature_flags + if context is None: + return False return context.get_bool( TelemetryHelper.TELEMETRY_FEATURE_FLAG_NAME, default_value=False ) diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py new file mode 100644 index 000000000..897d8774c --- /dev/null +++ b/tests/unit/conftest.py @@ -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 diff --git a/tests/unit/test_feature_flag_lifecycle.py b/tests/unit/test_feature_flag_lifecycle.py new file mode 100644 index 000000000..aedff39f9 --- /dev/null +++ b/tests/unit/test_feature_flag_lifecycle.py @@ -0,0 +1,97 @@ +from unittest.mock import MagicMock, patch + +import pytest + +from databricks import sql +from databricks.sql.backend.types import BackendType, SessionId +from databricks.sql.common.feature_flag import FeatureFlagsContextFactory +from databricks.sql.telemetry.telemetry_client import TelemetryHelper + + +@pytest.mark.parametrize("use_sea", [False, True]) +@pytest.mark.parametrize("oauth", [False, True]) +def test_flags_available_with_telemetry_disabled(session_feature_flags, use_sea, oauth): + session_feature_flags.get_instance.side_effect = ( + FeatureFlagsContextFactory.get_instance + ) + http = MagicMock() + http.request.return_value = MagicMock( + status=200, data=b'{"flags":[{"name":"sampleLimit","value":"42"}]}' + ) + auth = MagicMock() + auth.add_headers.side_effect = lambda headers: headers.update( + Authorization="Bearer test-token" + ) + backend = MagicMock() + backend.open_session.return_value = SessionId(BackendType.THRIFT, b"1", b"2") + + def create_backend(**kwargs): + http.request.assert_not_called() + assert kwargs["auth_provider"] is auth + return backend + + backend_name = "SeaDatabricksClient" if use_sea else "ThriftDatabricksClient" + try: + with patch("databricks.sql.client.UnifiedHttpClient", return_value=http), patch( + "databricks.sql.session.get_python_sql_connector_auth_provider", + return_value=auth, + ) as auth_factory, patch( + f"databricks.sql.session.{backend_name}", side_effect=create_backend + ): + conn = sql.connect( + "flags.example", + "/sql/1.0/warehouses/test?o=123", + enable_telemetry=False, + use_sea=use_sea, + **( + {"auth_type": "databricks-oauth"} + if oauth + else {"access_token": "test-token"} + ), + ) + try: + assert conn.telemetry_enabled is False + http.request.assert_not_called() + assert conn.session.feature_flags.get_int32("sampleLimit") == 42 + assert ( + http.request.call_args.kwargs["headers"]["Authorization"] + == "Bearer test-token" + ) + assert ( + http.request.call_args.kwargs["headers"]["x-databricks-org-id"] + == "123" + ) + conn.enable_telemetry = True + assert TelemetryHelper.is_telemetry_enabled(conn) is False + http.request.assert_called_once() + auth_factory.assert_called_once() + finally: + conn.close() + finally: + FeatureFlagsContextFactory.remove_instance( + "flags.example", {"x-databricks-org-id": "123"} + ) + + +def test_failed_flag_fetch_does_not_block_session(session_feature_flags): + session_feature_flags.get_instance.side_effect = ( + FeatureFlagsContextFactory.get_instance + ) + http = MagicMock() + http.request.side_effect = OSError("connector-service unavailable") + try: + with patch("databricks.sql.client.UnifiedHttpClient", return_value=http), patch( + "databricks.sql.session.ThriftDatabricksClient" + ) as backend: + backend.return_value.open_session.return_value = SessionId( + BackendType.THRIFT, b"1", b"2" + ) + with sql.connect( + "flags.example", "/test", access_token="test", enable_telemetry=False + ) as conn: + backend.return_value.open_session.assert_called_once() + http.request.assert_not_called() + assert conn.session.feature_flags.get_bool("missing") is False + http.request.assert_called_once() + finally: + FeatureFlagsContextFactory.remove_instance("flags.example") diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index 69522bc60..2d504735a 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -398,7 +398,9 @@ def _build_session(self, **extra): ) return sess, mock_get_provider - def test_use_kernel_m2m_does_not_build_connector_provider(self): + def test_use_kernel_m2m_does_not_build_connector_provider( + self, session_feature_flags + ): sess, mock_get_provider = self._build_session( oauth_client_id="sp-uuid", oauth_client_secret="shh" ) @@ -408,8 +410,12 @@ def test_use_kernel_m2m_does_not_build_connector_provider(self): # ...and with no access_token, auth_provider is None (M2M # resolves in-kernel from the raw kwargs). assert sess.auth_provider is None + assert sess.feature_flags is None + session_feature_flags.get_instance.assert_not_called() - def test_use_kernel_pat_builds_minimal_access_token_provider(self): + def test_use_kernel_pat_builds_minimal_access_token_provider( + self, session_feature_flags + ): from databricks.sql.auth.authenticators import AccessTokenAuthProvider sess, mock_get_provider = self._build_session(access_token="dapi-xyz") @@ -417,6 +423,8 @@ def test_use_kernel_pat_builds_minimal_access_token_provider(self): # PAT path: a minimal AccessTokenAuthProvider, not the # federation-wrapped connector provider. assert isinstance(sess.auth_provider, AccessTokenAuthProvider) + assert sess.feature_flags is None + session_feature_flags.get_instance.assert_not_called() class TestKernelTransportOptionsThreading: diff --git a/tests/unit/test_telemetry.py b/tests/unit/test_telemetry.py index 1beaf51e4..936b5fc5d 100644 --- a/tests/unit/test_telemetry.py +++ b/tests/unit/test_telemetry.py @@ -542,6 +542,11 @@ def teardown_method(self): TelemetryClientFactory._clients.clear() FeatureFlagsContextFactory._context_map.clear() + def _set_session_flags(self, session): + session.feature_flags = FeatureFlagsContextFactory.get_instance( + session.host, session.http_client, session.auth_provider, "test-agent" + ) + def _mock_ff_response(self, mock_http_request, enabled: bool): """Helper method to mock feature flag response for unified HTTP client.""" mock_response = MagicMock() @@ -577,6 +582,7 @@ def test_telemetry_enabled_when_flag_is_true(self, mock_http_request, MockSessio mock_http_client = MagicMock() mock_http_client.request = mock_http_request mock_session_instance.http_client = mock_http_client + self._set_session_flags(mock_session_instance) conn = sql.client.Connection( server_hostname="test", @@ -609,6 +615,7 @@ def test_telemetry_disabled_when_flag_is_false( mock_http_client = MagicMock() mock_http_client.request = mock_http_request mock_session_instance.http_client = mock_http_client + self._set_session_flags(mock_session_instance) conn = sql.client.Connection( server_hostname="test", @@ -641,6 +648,7 @@ def test_telemetry_disabled_when_flag_request_fails( mock_http_client = MagicMock() mock_http_client.request = mock_http_request mock_session_instance.http_client = mock_http_client + self._set_session_flags(mock_session_instance) conn = sql.client.Connection( server_hostname="test", From c2f065b750664548471d26d39a841a218479c1a1 Mon Sep 17 00:00:00 2001 From: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> Date: Wed, 7 Oct 2026 17:30:13 +0000 Subject: [PATCH 4/6] chore(kernel): bump to merged feature flag cache Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> --- KERNEL_REV | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/KERNEL_REV b/KERNEL_REV index f22f21685..037bffd30 100644 --- a/KERNEL_REV +++ b/KERNEL_REV @@ -1 +1 @@ -80f2aee7d884994d7b0af9a9ea6078872859a9cd +2c958bfba476a0b0f165c829988959a32bf12f3e From 103ff77769b4a9a21ec84c61944789def8f530a9 Mon Sep 17 00:00:00 2001 From: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> Date: Wed, 7 Oct 2026 18:06:18 +0000 Subject: [PATCH 5/6] fix: initialize session feature flags lazily Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> --- src/databricks/sql/session.py | 26 ++++++++++++++--------- tests/unit/test_feature_flag_lifecycle.py | 3 +++ 2 files changed, 19 insertions(+), 10 deletions(-) diff --git a/src/databricks/sql/session.py b/src/databricks/sql/session.py index eba208c61..74d1f4661 100644 --- a/src/databricks/sql/session.py +++ b/src/databricks/sql/session.py @@ -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 @@ -138,8 +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. - self.feature_flags: Optional[FeatureFlagsContext] = None - 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 @@ -148,13 +149,6 @@ 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, @@ -167,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, @@ -178,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( diff --git a/tests/unit/test_feature_flag_lifecycle.py b/tests/unit/test_feature_flag_lifecycle.py index aedff39f9..98dff5851 100644 --- a/tests/unit/test_feature_flag_lifecycle.py +++ b/tests/unit/test_feature_flag_lifecycle.py @@ -27,6 +27,7 @@ def test_flags_available_with_telemetry_disabled(session_feature_flags, use_sea, def create_backend(**kwargs): http.request.assert_not_called() + session_feature_flags.get_instance.assert_not_called() assert kwargs["auth_provider"] is auth return backend @@ -52,6 +53,7 @@ def create_backend(**kwargs): try: assert conn.telemetry_enabled is False http.request.assert_not_called() + session_feature_flags.get_instance.assert_not_called() assert conn.session.feature_flags.get_int32("sampleLimit") == 42 assert ( http.request.call_args.kwargs["headers"]["Authorization"] @@ -65,6 +67,7 @@ def create_backend(**kwargs): assert TelemetryHelper.is_telemetry_enabled(conn) is False http.request.assert_called_once() auth_factory.assert_called_once() + session_feature_flags.get_instance.assert_called_once() finally: conn.close() finally: From 8a2c905d7aa820b9f30bf703bbb7d27b7f82b5a4 Mon Sep 17 00:00:00 2001 From: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> Date: Wed, 7 Oct 2026 22:52:45 +0000 Subject: [PATCH 6/6] docs: clarify process-wide feature flag cache lifetime Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com> --- src/databricks/sql/common/feature_flag.py | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) diff --git a/src/databricks/sql/common/feature_flag.py b/src/databricks/sql/common/feature_flag.py index fb1fad527..a697e9da8 100644 --- a/src/databricks/sql/common/feature_flag.py +++ b/src/databricks/sql/common/feature_flag.py @@ -193,8 +193,10 @@ def _update_cache_from_response(self, ff_response: FeatureFlagsResponse): class FeatureFlagsContextFactory: """ - Shares flag values per workspace, independent of telemetry/session lifetime. - 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[tuple, _CacheState] = {} @@ -234,14 +236,17 @@ def get_instance( @classmethod def remove_instance(cls, host, headers=None): - """Evicts a workspace's values and shuts down the executor if the cache is empty.""" + """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: 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