Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .golangci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ linters:
- github.com/rs/zerolog
- github.com/stretchr/testify
- golang.org/x/oauth2
- golang.org/x/sync/singleflight
gosec:
severity: low
confidence: low
Expand Down
14 changes: 10 additions & 4 deletions connection.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,17 +20,19 @@ import (
"github.com/databricks/databricks-sql-go/internal/config"
"github.com/databricks/databricks-sql-go/internal/debuglog"
dbsqlerrint "github.com/databricks/databricks-sql-go/internal/errors"
"github.com/databricks/databricks-sql-go/internal/featureflags"
"github.com/databricks/databricks-sql-go/internal/retry"
"github.com/databricks/databricks-sql-go/internal/rows"
"github.com/databricks/databricks-sql-go/logger"
"github.com/databricks/databricks-sql-go/telemetry"
)

type conn struct {
id string
cfg *config.Config
backend backend.Backend
telemetry *telemetry.Interceptor // Optional telemetry interceptor
id string
cfg *config.Config
backend backend.Backend
featureFlags *featureflags.Request
telemetry *telemetry.Interceptor // Optional telemetry interceptor
}

// tagStatementClosed returns a telemetry-only copy of a close-RPC error tagged
Expand Down Expand Up @@ -77,6 +79,10 @@ func (c *conn) Close() error {
_ = c.telemetry.Close(ctx)
telemetry.ReleaseForConnection(c.cfg.Host)
}
if c.featureFlags != nil {
featureflags.GetCache().Release(c.featureFlags.Host, c.featureFlags.WorkspaceID)
c.featureFlags = nil
}

if err != nil {
log.Err(err).Msg("databricks: failed to close connection")
Expand Down
40 changes: 28 additions & 12 deletions connector.go
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ import (
"github.com/databricks/databricks-sql-go/internal/client"
"github.com/databricks/databricks-sql-go/internal/config"
"github.com/databricks/databricks-sql-go/internal/debuglog"
"github.com/databricks/databricks-sql-go/internal/featureflags"
"github.com/databricks/databricks-sql-go/internal/warehouse_cache"
"github.com/databricks/databricks-sql-go/logger"
"github.com/databricks/databricks-sql-go/telemetry"
Expand Down Expand Up @@ -66,11 +67,19 @@ type federatedTokenAuthenticator struct {
func (c *connector) Connect(ctx context.Context) (driver.Conn, error) {
defer debuglog.Track(ctx, "connector.Connect", "host=%s", c.cfg.Host)()

flagRequest := c.featureFlagRequest()
if !c.cfg.UseKernel {
featureflags.GetCache().Acquire(flagRequest.Host, flagRequest.WorkspaceID)
}

// openSessionWithReydenFallback handles the session opening with automatic
// recovery for Reyden / Real-Time warehouses that reject Thrift. It returns
// the backend, latency, and error.
be, sessionLatencyMs, err := c.openSessionWithReydenFallback(ctx)
if err != nil {
if !c.cfg.UseKernel {
featureflags.GetCache().Release(flagRequest.Host, flagRequest.WorkspaceID)
}
return nil, err
}

Expand All @@ -79,17 +88,10 @@ func (c *connector) Connect(ctx context.Context) (driver.Conn, error) {
cfg: c.cfg,
backend: be,
}
log := logger.WithContext(conn.id, driverctx.CorrelationIdFromContext(ctx), "")

// Extract SPOG routing headers from HTTPPath. When the workspace ID is
// available via ?o=<workspaceId> or a cluster /o/<workspaceId>/ path segment,
// wrap the HTTP client used for telemetry + feature-flag calls with a
// transport that injects x-databricks-org-id. Thrift routes via the URL so
// its own c.client doesn't need wrapping.
telemetryClient := c.client
if spogHeaders := extractSpogHeaders(c.cfg.HTTPPath); len(spogHeaders) > 0 {
telemetryClient = withSpogHeaders(c.client, spogHeaders)
if !c.cfg.UseKernel {
conn.featureFlags = &flagRequest
}
log := logger.WithContext(conn.id, driverctx.CorrelationIdFromContext(ctx), "")

// Skip driver telemetry on the kernel path. The kernel owns query execution
// below the driver backend, so keeping the Go telemetry interceptor active
Expand All @@ -103,9 +105,10 @@ func (c *connector) Connect(ctx context.Context) (driver.Conn, error) {
if !skipTelemetry {
conn.telemetry = telemetry.InitializeForConnection(ctx, telemetry.TelemetryInitOptions{
Host: c.cfg.Host,
WorkspaceID: flagRequest.WorkspaceID,
DriverVersion: c.cfg.DriverVersion,
UserAgent: client.BuildUserAgent(c.cfg),
HTTPClient: telemetryClient,
UserAgent: flagRequest.UserAgent,
HTTPClient: flagRequest.HTTPClient,
EnableTelemetry: c.cfg.EnableTelemetry,
BatchSize: c.cfg.TelemetryBatchSize,
FlushInterval: c.cfg.TelemetryFlushInterval,
Expand All @@ -128,6 +131,19 @@ func (c *connector) Connect(ctx context.Context) (driver.Conn, error) {
return conn, nil
}

func (c *connector) featureFlagRequest() featureflags.Request {
headers := extractSpogHeaders(c.cfg.HTTPPath)
httpClient := c.client
if len(headers) > 0 {
httpClient = withSpogHeaders(httpClient, headers)
}
return featureflags.Request{
Host: c.cfg.Host, WorkspaceID: headers["x-databricks-org-id"],
DriverVersion: c.cfg.DriverVersion, UserAgent: client.BuildUserAgent(c.cfg),
HTTPClient: httpClient,
}
}

// Driver returns underlying databricksDriver for compatibility with sql.DB Driver method
func (c *connector) Driver() driver.Driver {
return &databricksDriver{}
Expand Down
126 changes: 126 additions & 0 deletions connector_feature_flags_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,126 @@
package dbsql

import (
"context"
"errors"
"io"
"net/http"
"strings"
"sync/atomic"
"testing"

"github.com/databricks/databricks-sql-go/internal/backend"
"github.com/databricks/databricks-sql-go/internal/backend/thrift"
"github.com/databricks/databricks-sql-go/internal/cli_service"
"github.com/databricks/databricks-sql-go/internal/client"
"github.com/databricks/databricks-sql-go/internal/config"
"github.com/databricks/databricks-sql-go/internal/featureflags"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

type flagRoundTripper func(*http.Request) (*http.Response, error)

func (f flagRoundTripper) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) }

func TestConnectionOwnsFlagsWithTelemetryDisabled(t *testing.T) {
var fetches atomic.Int32
transport := flagRoundTripper(func(r *http.Request) (*http.Response, error) {
fetches.Add(1)
assert.Equal(t, "Bearer test-token", r.Header.Get("Authorization"))
assert.Equal(t, "123", r.Header.Get("X-Databricks-Org-Id"))
assert.Equal(t, "/api/2.0/connector-service/feature-flags/GOLANG/"+DriverVersion, r.URL.Path)
return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(
`{"flags":[{"name":"sampleLimit","value":"42"}]}`)), Header: make(http.Header)}, nil
})
driverConnector, err := NewConnector(
WithServerHostname("flags.example"), WithHTTPPath("/sql/1.0/warehouses/test?o=123"),
WithAccessToken("test-token"), WithTransport(transport),
func(cfg *config.Config) { cfg.EnableTelemetry = config.NewConfigValue(false) },
)
require.NoError(t, err)
c := driverConnector.(*connector)
ctx := context.Background()
flags := featureflags.GetCache()
request := c.featureFlagRequest()
c.thriftBackendFactory = func(ctx context.Context, cfg *config.Config, _ *http.Client) (backend.Backend, error) {
return thrift.NewForTest(&client.TestClient{
FnOpenSession: func(context.Context, *cli_service.TOpenSessionReq) (*cli_service.TOpenSessionResp, error) {
return getTestSession(), nil
},
FnCloseSession: func(context.Context, *cli_service.TCloseSessionReq) (*cli_service.TCloseSessionResp, error) {
return &cli_service.TCloseSessionResp{Status: &cli_service.TStatus{StatusCode: cli_service.TStatusCode_SUCCESS_STATUS}}, nil
},
}, getTestSession(), cfg), nil
}
first, err := c.Connect(ctx)
require.NoError(t, err)
second, err := c.Connect(ctx)
require.NoError(t, err)
require.Nil(t, first.(*conn).telemetry)
require.Zero(t, fetches.Load(), "unused flags must not add connection requests")
value, err := flags.GetInt32(ctx, *first.(*conn).featureFlags, "sampleLimit")
require.NoError(t, err)
require.EqualValues(t, 42, value)
require.NoError(t, first.Close())
value, err = flags.GetInt32(ctx, *second.(*conn).featureFlags, "sampleLimit")
require.NoError(t, err)
require.EqualValues(t, 42, value)
require.EqualValues(t, 1, fetches.Load())
require.NoError(t, second.Close())
flags.Acquire(request.Host, request.WorkspaceID)
defer flags.Release(request.Host, request.WorkspaceID)
_, err = flags.GetInt32(ctx, request, "sampleLimit")
require.NoError(t, err)
require.EqualValues(t, 2, fetches.Load(), "last close releases the workspace cache")
}

func TestConnectionFlagFailureAndKernelPaths(t *testing.T) {
for _, scenario := range []string{"flag fetch fails", "backend fails", "explicit kernel"} {
t.Run(scenario, func(t *testing.T) {
var fetches atomic.Int32
c, err := NewConnector(WithServerHostname("flags.example"), WithAccessToken("test-token"),
WithTransport(flagRoundTripper(func(*http.Request) (*http.Response, error) {
fetches.Add(1)
status := http.StatusOK
if scenario == "flag fetch fails" {
status = http.StatusServiceUnavailable
}
return &http.Response{StatusCode: status, Body: io.NopCloser(strings.NewReader(`{"flags":[]}`)), Header: make(http.Header)}, nil
})), WithUseKernel(scenario == "explicit kernel"))
require.NoError(t, err)
connector := c.(*connector)
connector.thriftBackendFactory = func(context.Context, *config.Config, *http.Client) (backend.Backend, error) {
if scenario == "backend fails" {
return nil, errors.New("backend failed")
}
return &fakeThriftBackend{sessionID: "test-session"}, nil
}
connector.kernelBackendFactory = func(context.Context, *config.Config) (backend.Backend, error) {
return &fakeKernelBackend{}, nil
}
connection, err := c.Connect(context.Background())
if scenario == "backend fails" {
require.ErrorContains(t, err, "backend failed")
} else {
require.NoError(t, err)
require.Zero(t, fetches.Load())
if scenario == "flag fetch fails" {
value, err := featureflags.GetCache().GetBool(context.Background(), *connection.(*conn).featureFlags, "missing")
require.Error(t, err)
require.False(t, value)
}
require.NoError(t, connection.Close())
}
// No consumer is left after a failed open or close; a getter cannot fetch.
value, err := featureflags.GetCache().GetBool(context.Background(), connector.featureFlagRequest(), "missing")
require.NoError(t, err)
require.False(t, value)
if scenario == "flag fetch fails" {
require.EqualValues(t, 1, fetches.Load())
} else {
require.Zero(t, fetches.Load())
}
})
}
}
2 changes: 1 addition & 1 deletion go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ require (
github.com/pierrec/lz4/v4 v4.1.15
github.com/stretchr/testify v1.8.1
golang.org/x/oauth2 v0.27.0
golang.org/x/sync v0.22.0
gotest.tools/gotestsum v1.8.2
)

Expand All @@ -38,7 +39,6 @@ require (
github.com/zeebo/xxh3 v1.0.2 // indirect
golang.org/x/crypto v0.55.0 // indirect
golang.org/x/mod v0.40.0 // indirect
golang.org/x/sync v0.22.0 // indirect
golang.org/x/telemetry v0.0.0-20260811182544-a038080d80e5 // indirect
golang.org/x/term v0.45.0 // indirect
golang.org/x/tools v0.49.0 // indirect
Expand Down
Loading
Loading