11import json
2+ import math
23import threading
34import time
45from 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
89from databricks .sql .common .http import HttpMethod
910from 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
1614class FeatureFlagEntry :
@@ -43,77 +41,126 @@ def from_dict(cls, data: Dict[str, Any]) -> "FeatureFlagsResponse":
4341REFRESH_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+
4663class 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
146193class 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
0 commit comments