Skip to content

Commit 1ce7f5b

Browse files
committed
Generalize driver feature flag cache and typed reads
Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com>
1 parent 21d288f commit 1ce7f5b

3 files changed

Lines changed: 216 additions & 66 deletions

File tree

‎src/databricks/sql/common/feature_flag.py‎

Lines changed: 111 additions & 57 deletions
Original file line numberDiff line numberDiff line change
@@ -1,16 +1,14 @@
11
import json
2+
import math
23
import threading
34
import time
45
from dataclasses import dataclass, field
5-
from concurrent.futures import ThreadPoolExecutor
6-
from typing import Dict, Optional, List, Any, TYPE_CHECKING
6+
from concurrent.futures import Future, ThreadPoolExecutor
7+
from typing import Dict, Optional, List, Any
78

89
from databricks.sql.common.http import HttpMethod
910
from databricks.sql.common.url_utils import normalize_host_with_protocol
1011

11-
if TYPE_CHECKING:
12-
from databricks.sql.client import Connection
13-
1412

1513
@dataclass
1614
class FeatureFlagEntry:
@@ -43,77 +41,126 @@ def from_dict(cls, data: Dict[str, Any]) -> "FeatureFlagsResponse":
4341
REFRESH_BEFORE_EXPIRY_SECONDS = 10 # Start proactive refresh 10s before expiry
4442

4543

44+
@dataclass
45+
class _CacheState:
46+
# Only values/coordination are shared; credentials and HTTP clients are not.
47+
flags: Optional[Dict[str, str]] = None
48+
ttl_seconds: int = DEFAULT_TTL_SECONDS
49+
last_refresh_time: float = 0
50+
lock: Any = field(default_factory=threading.RLock)
51+
refresh: Optional[Future] = None
52+
53+
54+
def _cache_key(host, headers):
55+
workspace_id = (headers or {}).get("x-databricks-org-id")
56+
return (
57+
("workspace", workspace_id)
58+
if workspace_id
59+
else ("host", normalize_host_with_protocol(host).lower())
60+
)
61+
62+
4663
class FeatureFlagsContext:
4764
"""
48-
Manages fetching and caching of server-side feature flags for a connection.
65+
Authenticated flag reader usable before any session/backend is opened.
4966
5067
1. The very first check for any flag is a synchronous, BLOCKING operation.
5168
2. Subsequent refreshes (triggered near TTL expiry) are done asynchronously
5269
in the background, returning stale data until the refresh completes.
5370
"""
5471

5572
def __init__(
56-
self, connection: "Connection", executor: ThreadPoolExecutor, http_client
73+
self, host, executor, http_client, auth_provider, user_agent, headers, state
5774
):
5875
from databricks.sql import __version__
5976

60-
self._connection = connection
6177
self._executor = executor # Used for ASYNCHRONOUS refreshes
62-
self._lock = threading.RLock()
63-
64-
# Cache state: `None` indicates the cache has never been loaded.
65-
self._flags: Optional[Dict[str, str]] = None
66-
self._ttl_seconds: int = DEFAULT_TTL_SECONDS
67-
self._last_refresh_time: float = 0
78+
self._state = state
79+
self._auth_provider = auth_provider
80+
self._headers = {"User-Agent": user_agent, **headers}
6881

6982
endpoint_suffix = FEATURE_FLAGS_ENDPOINT_SUFFIX_FORMAT.format(__version__)
7083
self._feature_flag_endpoint = (
71-
normalize_host_with_protocol(self._connection.session.host)
72-
+ endpoint_suffix
84+
normalize_host_with_protocol(host) + endpoint_suffix
7385
)
7486

7587
# Use the provided HTTP client
7688
self._http_client = http_client
7789

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

83-
refresh_threshold = self._last_refresh_time + (
84-
self._ttl_seconds - REFRESH_BEFORE_EXPIRY_SECONDS
95+
refresh_threshold = self._state.last_refresh_time + (
96+
self._state.ttl_seconds - REFRESH_BEFORE_EXPIRY_SECONDS
8597
)
8698
return time.monotonic() > refresh_threshold
8799

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

99111
# If a proactive background refresh is needed, start one. This is non-blocking.
100-
elif self._is_refresh_needed():
101-
# We don't check for an in-flight refresh; the executor queues the task, which is safe.
102-
self._executor.submit(self._refresh_flags)
112+
elif self._is_refresh_needed() and (
113+
self._state.refresh is None or self._state.refresh.done()
114+
):
115+
self._state.refresh = self._executor.submit(self._refresh_flags)
103116

104-
assert self._flags is not None
117+
raw = (self._state.flags or {}).get(name)
118+
try:
119+
return json.loads(raw) if raw is not None else None
120+
except (TypeError, ValueError):
121+
return None
122+
123+
def get_bool(self, name: str, default_value: bool = False) -> bool:
124+
value = self._get_value(name)
125+
return value if type(value) is bool else default_value
126+
127+
def _get_int(self, name, bits, default_value):
128+
value = self._get_value(name)
129+
if type(value) is int and -(2 ** (bits - 1)) <= value < 2 ** (bits - 1):
130+
return value
131+
return default_value
132+
133+
def get_int32(self, name: str, default_value=None) -> Optional[int]:
134+
return self._get_int(name, 32, default_value)
105135

106-
# Now, return the value from the populated cache.
107-
return self._flags.get(name, default_value)
136+
def get_int64(self, name: str, default_value=None) -> Optional[int]:
137+
return self._get_int(name, 64, default_value)
138+
139+
def get_double(self, name: str, default_value=None) -> Optional[float]:
140+
value = self._get_value(name)
141+
try:
142+
if type(value) in (int, float) and math.isfinite(value):
143+
return float(value)
144+
except OverflowError:
145+
pass
146+
return default_value
147+
148+
def get_string(self, name: str, default_value=None) -> Optional[str]:
149+
value = self._get_value(name)
150+
return value if isinstance(value, str) else default_value
151+
152+
def get_string_list(self, name: str, default_value=None) -> Optional[List[str]]:
153+
value = self._get_value(name)
154+
if isinstance(value, list) and all(isinstance(item, str) for item in value):
155+
return value
156+
return default_value
108157

109158
def _refresh_flags(self):
110159
"""Performs a synchronous network request to fetch and update flags."""
111-
headers = {}
160+
headers = dict(self._headers)
112161
try:
113162
# Authenticate the request
114-
self._connection.session.auth_provider.add_headers(headers)
115-
headers["User-Agent"] = self._connection.session.useragent_header
116-
headers.update(self._connection.session.get_spog_headers())
163+
self._auth_provider.add_headers(headers)
117164

118165
response = self._http_client.request(
119166
HttpMethod.GET, self._feature_flag_endpoint, headers=headers, timeout=30
@@ -126,30 +173,30 @@ def _refresh_flags(self):
126173
self._update_cache_from_response(ff_response)
127174
else:
128175
# On failure, initialize with an empty dictionary to prevent re-blocking.
129-
if self._flags is None:
130-
self._flags = {}
176+
if self._state.flags is None:
177+
self._state.flags = {}
131178

132-
except Exception as e:
179+
except Exception:
133180
# On exception, initialize with an empty dictionary to prevent re-blocking.
134-
if self._flags is None:
135-
self._flags = {}
181+
if self._state.flags is None:
182+
self._state.flags = {}
136183

137184
def _update_cache_from_response(self, ff_response: FeatureFlagsResponse):
138185
"""Atomically updates the internal cache state from a successful server response."""
139-
with self._lock:
140-
self._flags = {flag.name: flag.value for flag in ff_response.flags}
186+
with self._state.lock:
187+
self._state.flags = {flag.name: flag.value for flag in ff_response.flags}
141188
if ff_response.ttl_seconds is not None and ff_response.ttl_seconds > 0:
142-
self._ttl_seconds = ff_response.ttl_seconds
143-
self._last_refresh_time = time.monotonic()
189+
self._state.ttl_seconds = ff_response.ttl_seconds
190+
self._state.last_refresh_time = time.monotonic()
144191

145192

146193
class FeatureFlagsContextFactory:
147194
"""
148-
Manages a singleton instance of FeatureFlagsContext per connection session.
195+
Shares flag values per workspace, independent of telemetry/session lifetime.
149196
Also manages a shared ThreadPoolExecutor for all background refresh operations.
150197
"""
151198

152-
_context_map: Dict[str, FeatureFlagsContext] = {}
199+
_context_map: Dict[tuple, _CacheState] = {}
153200
_executor: Optional[ThreadPoolExecutor] = None
154201
_lock = threading.Lock()
155202

@@ -162,27 +209,34 @@ def _initialize(cls):
162209
)
163210

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

171-
# Cache at HOST level - share feature flags across connections to same host
172-
# Feature flags are per-host, not per-session
173-
key = connection.session.host
221+
key = _cache_key(host, headers)
174222
if key not in cls._context_map:
175-
cls._context_map[key] = FeatureFlagsContext(
176-
connection, cls._executor, connection.session.http_client
177-
)
178-
return cls._context_map[key]
223+
cls._context_map[key] = _CacheState()
224+
return FeatureFlagsContext(
225+
host,
226+
cls._executor,
227+
http_client,
228+
auth_provider,
229+
user_agent,
230+
headers,
231+
cls._context_map[key],
232+
)
179233

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

‎src/databricks/sql/telemetry/telemetry_client.py‎

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -134,11 +134,17 @@ def is_telemetry_enabled(connection: "Connection") -> bool:
134134
return False
135135

136136
# Only fetch feature flags when enable_telemetry=True and not forced
137-
context = FeatureFlagsContextFactory.get_instance(connection)
138-
flag_value = context.get_flag_value(
137+
session = connection.session
138+
context = FeatureFlagsContextFactory.get_instance(
139+
session.host,
140+
session.http_client,
141+
session.auth_provider,
142+
session.useragent_header,
143+
session.get_spog_headers(),
144+
)
145+
return context.get_bool(
139146
TelemetryHelper.TELEMETRY_FEATURE_FLAG_NAME, default_value=False
140147
)
141-
return str(flag_value).lower() == "true"
142148

143149

144150
class NoopTelemetryClient(BaseTelemetryClient):

0 commit comments

Comments
 (0)