Skip to content

Commit 9fda19d

Browse files
committed
fix(feature-flags): decouple cache setup from telemetry
Signed-off-by: Cathleen Yan <58714163+cathleeny@users.noreply.github.com>
1 parent 0ea879a commit 9fda19d

6 files changed

Lines changed: 142 additions & 12 deletions

File tree

‎src/databricks/sql/session.py‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,10 @@
1313
from databricks.sql.backend.types import SessionId, BackendType
1414
from databricks.sql.common.unified_http_client import UnifiedHttpClient
1515
from databricks.sql.common.agent import detect as detect_agent
16+
from databricks.sql.common.feature_flag import (
17+
FeatureFlagsContext,
18+
FeatureFlagsContextFactory,
19+
)
1620
from databricks.sql.telemetry.telemetry_client import TelemetryClientFactory
1721

1822
if TYPE_CHECKING:
@@ -134,6 +138,7 @@ def __init__(
134138
# provider when an ``access_token`` is present, and ``None``
135139
# otherwise (OAuth M2M/U2M resolve purely from the raw kwargs
136140
# the bridge reads). The Thrift / SEA backends are unchanged.
141+
self.feature_flags: Optional[FeatureFlagsContext] = None
137142
if kwargs.get("use_kernel", False):
138143
access_token = kwargs.get("access_token")
139144
self.auth_provider = (
@@ -143,6 +148,13 @@ def __init__(
143148
self.auth_provider = get_python_sql_connector_auth_provider(
144149
server_hostname, http_client=self.http_client, **kwargs
145150
)
151+
self.feature_flags = FeatureFlagsContextFactory.get_instance(
152+
self.host,
153+
self.http_client,
154+
self.auth_provider,
155+
self.useragent_header,
156+
self.get_spog_headers(),
157+
)
146158

147159
self.backend = self._create_backend(
148160
server_hostname,

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

Lines changed: 3 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -40,7 +40,6 @@
4040
import uuid
4141
import locale
4242
from databricks.sql.telemetry.utils import BaseTelemetryClient
43-
from databricks.sql.common.feature_flag import FeatureFlagsContextFactory
4443
from databricks.sql.common.unified_http_client import UnifiedHttpClient
4544
from databricks.sql.common.http import HttpMethod
4645
from databricks.sql.exc import RequestError
@@ -133,15 +132,9 @@ def is_telemetry_enabled(connection: "Connection") -> bool:
133132
if not connection.enable_telemetry:
134133
return False
135134

136-
# Only fetch feature flags when enable_telemetry=True and not forced
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-
)
135+
context = connection.session.feature_flags
136+
if context is None:
137+
return False
145138
return context.get_bool(
146139
TelemetryHelper.TELEMETRY_FEATURE_FLAG_NAME, default_value=False
147140
)

‎tests/unit/conftest.py‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,12 @@
1+
from unittest.mock import patch
2+
3+
import pytest
4+
5+
6+
@pytest.fixture(autouse=True)
7+
def session_feature_flags():
8+
# Unit connections must not contact connector-service. Cache/lifecycle tests
9+
# exercise the real reader with a mocked HTTP client.
10+
with patch("databricks.sql.session.FeatureFlagsContextFactory") as factory:
11+
factory.get_instance.return_value.get_bool.return_value = False
12+
yield factory
Lines changed: 97 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,97 @@
1+
from unittest.mock import MagicMock, patch
2+
3+
import pytest
4+
5+
from databricks import sql
6+
from databricks.sql.backend.types import BackendType, SessionId
7+
from databricks.sql.common.feature_flag import FeatureFlagsContextFactory
8+
from databricks.sql.telemetry.telemetry_client import TelemetryHelper
9+
10+
11+
@pytest.mark.parametrize("use_sea", [False, True])
12+
@pytest.mark.parametrize("oauth", [False, True])
13+
def test_flags_available_with_telemetry_disabled(session_feature_flags, use_sea, oauth):
14+
session_feature_flags.get_instance.side_effect = (
15+
FeatureFlagsContextFactory.get_instance
16+
)
17+
http = MagicMock()
18+
http.request.return_value = MagicMock(
19+
status=200, data=b'{"flags":[{"name":"sampleLimit","value":"42"}]}'
20+
)
21+
auth = MagicMock()
22+
auth.add_headers.side_effect = lambda headers: headers.update(
23+
Authorization="Bearer test-token"
24+
)
25+
backend = MagicMock()
26+
backend.open_session.return_value = SessionId(BackendType.THRIFT, b"1", b"2")
27+
28+
def create_backend(**kwargs):
29+
http.request.assert_not_called()
30+
assert kwargs["auth_provider"] is auth
31+
return backend
32+
33+
backend_name = "SeaDatabricksClient" if use_sea else "ThriftDatabricksClient"
34+
try:
35+
with patch("databricks.sql.client.UnifiedHttpClient", return_value=http), patch(
36+
"databricks.sql.session.get_python_sql_connector_auth_provider",
37+
return_value=auth,
38+
) as auth_factory, patch(
39+
f"databricks.sql.session.{backend_name}", side_effect=create_backend
40+
):
41+
conn = sql.connect(
42+
"flags.example",
43+
"/sql/1.0/warehouses/test?o=123",
44+
enable_telemetry=False,
45+
use_sea=use_sea,
46+
**(
47+
{"auth_type": "databricks-oauth"}
48+
if oauth
49+
else {"access_token": "test-token"}
50+
),
51+
)
52+
try:
53+
assert conn.telemetry_enabled is False
54+
http.request.assert_not_called()
55+
assert conn.session.feature_flags.get_int32("sampleLimit") == 42
56+
assert (
57+
http.request.call_args.kwargs["headers"]["Authorization"]
58+
== "Bearer test-token"
59+
)
60+
assert (
61+
http.request.call_args.kwargs["headers"]["x-databricks-org-id"]
62+
== "123"
63+
)
64+
conn.enable_telemetry = True
65+
assert TelemetryHelper.is_telemetry_enabled(conn) is False
66+
http.request.assert_called_once()
67+
auth_factory.assert_called_once()
68+
finally:
69+
conn.close()
70+
finally:
71+
FeatureFlagsContextFactory.remove_instance(
72+
"flags.example", {"x-databricks-org-id": "123"}
73+
)
74+
75+
76+
def test_failed_flag_fetch_does_not_block_session(session_feature_flags):
77+
session_feature_flags.get_instance.side_effect = (
78+
FeatureFlagsContextFactory.get_instance
79+
)
80+
http = MagicMock()
81+
http.request.side_effect = OSError("connector-service unavailable")
82+
try:
83+
with patch("databricks.sql.client.UnifiedHttpClient", return_value=http), patch(
84+
"databricks.sql.session.ThriftDatabricksClient"
85+
) as backend:
86+
backend.return_value.open_session.return_value = SessionId(
87+
BackendType.THRIFT, b"1", b"2"
88+
)
89+
with sql.connect(
90+
"flags.example", "/test", access_token="test", enable_telemetry=False
91+
) as conn:
92+
backend.return_value.open_session.assert_called_once()
93+
http.request.assert_not_called()
94+
assert conn.session.feature_flags.get_bool("missing") is False
95+
http.request.assert_called_once()
96+
finally:
97+
FeatureFlagsContextFactory.remove_instance("flags.example")

‎tests/unit/test_session.py‎

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -398,7 +398,9 @@ def _build_session(self, **extra):
398398
)
399399
return sess, mock_get_provider
400400

401-
def test_use_kernel_m2m_does_not_build_connector_provider(self):
401+
def test_use_kernel_m2m_does_not_build_connector_provider(
402+
self, session_feature_flags
403+
):
402404
sess, mock_get_provider = self._build_session(
403405
oauth_client_id="sp-uuid", oauth_client_secret="shh"
404406
)
@@ -408,15 +410,21 @@ def test_use_kernel_m2m_does_not_build_connector_provider(self):
408410
# ...and with no access_token, auth_provider is None (M2M
409411
# resolves in-kernel from the raw kwargs).
410412
assert sess.auth_provider is None
413+
assert sess.feature_flags is None
414+
session_feature_flags.get_instance.assert_not_called()
411415

412-
def test_use_kernel_pat_builds_minimal_access_token_provider(self):
416+
def test_use_kernel_pat_builds_minimal_access_token_provider(
417+
self, session_feature_flags
418+
):
413419
from databricks.sql.auth.authenticators import AccessTokenAuthProvider
414420

415421
sess, mock_get_provider = self._build_session(access_token="dapi-xyz")
416422
mock_get_provider.assert_not_called()
417423
# PAT path: a minimal AccessTokenAuthProvider, not the
418424
# federation-wrapped connector provider.
419425
assert isinstance(sess.auth_provider, AccessTokenAuthProvider)
426+
assert sess.feature_flags is None
427+
session_feature_flags.get_instance.assert_not_called()
420428

421429

422430
class TestKernelTransportOptionsThreading:

‎tests/unit/test_telemetry.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -542,6 +542,11 @@ def teardown_method(self):
542542
TelemetryClientFactory._clients.clear()
543543
FeatureFlagsContextFactory._context_map.clear()
544544

545+
def _set_session_flags(self, session):
546+
session.feature_flags = FeatureFlagsContextFactory.get_instance(
547+
session.host, session.http_client, session.auth_provider, "test-agent"
548+
)
549+
545550
def _mock_ff_response(self, mock_http_request, enabled: bool):
546551
"""Helper method to mock feature flag response for unified HTTP client."""
547552
mock_response = MagicMock()
@@ -577,6 +582,7 @@ def test_telemetry_enabled_when_flag_is_true(self, mock_http_request, MockSessio
577582
mock_http_client = MagicMock()
578583
mock_http_client.request = mock_http_request
579584
mock_session_instance.http_client = mock_http_client
585+
self._set_session_flags(mock_session_instance)
580586

581587
conn = sql.client.Connection(
582588
server_hostname="test",
@@ -609,6 +615,7 @@ def test_telemetry_disabled_when_flag_is_false(
609615
mock_http_client = MagicMock()
610616
mock_http_client.request = mock_http_request
611617
mock_session_instance.http_client = mock_http_client
618+
self._set_session_flags(mock_session_instance)
612619

613620
conn = sql.client.Connection(
614621
server_hostname="test",
@@ -641,6 +648,7 @@ def test_telemetry_disabled_when_flag_request_fails(
641648
mock_http_client = MagicMock()
642649
mock_http_client.request = mock_http_request
643650
mock_session_instance.http_client = mock_http_client
651+
self._set_session_flags(mock_session_instance)
644652

645653
conn = sql.client.Connection(
646654
server_hostname="test",

0 commit comments

Comments
 (0)