diff --git a/src/databricks/sql/common/feature_flag.py b/src/databricks/sql/common/feature_flag.py index 0b2c7490b..fb1fad527 100644 --- a/src/databricks/sql/common/feature_flag.py +++ b/src/databricks/sql/common/feature_flag.py @@ -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: @@ -43,9 +42,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 +71,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 +90,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) - - 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 @@ -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() @@ -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) 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()