Skip to content

Commit 7ff28cb

Browse files
author
Akangsha Goel
committed
refactor(sources): reach connection handles through ConnectOnce
Every source builds its handle through sources.ConnectOnce and reaches it through an accessor, instead of connecting inline in Initialize and storing the result on an exported field. Initialize takes a deferConnect parameter, but the only caller passes false, so each source still connects during startup and still reports a connect failure there. The flag that sets it lands separately. Config that needs no network is resolved in newSource, which runs whether or not the connect is deferred: a malformed queryTimeout, baseUrl or writeMode is a configuration error and must fail at startup rather than on the first tool call. Sources keep their context-free accessors, which report the handle only once connected. They exist so tools can express a capability as an interface and type-assert on it; none of them is invoked. The exception is Looker, where LookerApiSettings is read after GetLookerSDK has connected. ConnectOnce bounds the attempt at ConnectTimeout, which the startup connect did not have before. Sources whose own config permits a longer connect raise the ceiling to match, as postgres and looker already did: mysql, oceanbase, singlestore and mindsdb through the driver read timeout the ping honours, trino through the timeout it sends with that query, and cockroachdb through the backoff its retry loop sleeps. The branch was cut before #3902, #3905 and #3921 landed, so it also restores alloydbpg's pre-PG17 read-only diagnostic and three test files it would otherwise have reverted, detaches the bigquery client creator from the connect's context, and stops InitConnectionSpan panicking on the nil tracer that some source tests still pass.
1 parent 0001190 commit 7ff28cb

63 files changed

Lines changed: 3376 additions & 1486 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎internal/server/primitives/primitives_test.go‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@ import (
1818
"testing"
1919

2020
"github.com/google/go-cmp/cmp"
21+
"github.com/google/go-cmp/cmp/cmpopts"
2122
"github.com/googleapis/mcp-toolbox/internal/auth"
2223
"github.com/googleapis/mcp-toolbox/internal/embeddingmodels"
2324
"github.com/googleapis/mcp-toolbox/internal/group"
@@ -48,7 +49,7 @@ func TestUpdateServer(t *testing.T) {
4849
resMgr := primitives.NewPrimitiveManager(newSources, newAuth, newEmbeddingModels, newTools, newPrompts, newGroups)
4950

5051
gotSource, _ := resMgr.GetSource("example-source")
51-
if diff := cmp.Diff(gotSource, newSources["example-source"]); diff != "" {
52+
if diff := cmp.Diff(gotSource, newSources["example-source"], cmpopts.IgnoreUnexported(alloydbpg.Source{})); diff != "" {
5253
t.Errorf("error updating server, sources (-want +got):\n%s", diff)
5354
}
5455

@@ -87,7 +88,7 @@ func TestUpdateServer(t *testing.T) {
8788

8889
resMgr.SetPrimitives(updateSource, newAuth, newEmbeddingModels, newTools, newPrompts, newGroups)
8990
gotSource, _ = resMgr.GetSource("example-source2")
90-
if diff := cmp.Diff(gotSource, updateSource["example-source2"]); diff != "" {
91+
if diff := cmp.Diff(gotSource, updateSource["example-source2"], cmpopts.IgnoreUnexported(alloydbpg.Source{})); diff != "" {
9192
t.Errorf("error updating server, sources (-want +got):\n%s", diff)
9293
}
9394
}

‎internal/server/server.go‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -113,7 +113,8 @@ func InitializeConfigs(ctx context.Context, cfg ServerConfig) (
113113
trace.WithAttributes(attribute.String("source_name", name)),
114114
)
115115
defer span.End()
116-
s, err := sc.Initialize(childCtx, instrumentation.Tracer)
116+
// Always connects here; the flag that defers it lands separately.
117+
s, err := sc.Initialize(childCtx, instrumentation.Tracer, false)
117118
if err != nil {
118119
return nil, fmt.Errorf("unable to initialize source %q: %w", name, err)
119120
}

‎internal/server/server_test.go‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@ import (
3838
"time"
3939

4040
"github.com/google/go-cmp/cmp"
41+
"github.com/google/go-cmp/cmp/cmpopts"
4142
"github.com/googleapis/mcp-toolbox/internal/auth"
4243
"github.com/googleapis/mcp-toolbox/internal/auth/generic"
4344
"github.com/googleapis/mcp-toolbox/internal/embeddingmodels"
@@ -432,7 +433,7 @@ func TestUpdateServer(t *testing.T) {
432433
}
433434

434435
gotSource, _ := s.PrimitiveMgr.GetSource("example-source")
435-
if diff := cmp.Diff(gotSource, newSources["example-source"]); diff != "" {
436+
if diff := cmp.Diff(gotSource, newSources["example-source"], cmpopts.IgnoreUnexported(alloydbpg.Source{})); diff != "" {
436437
t.Errorf("error updating server, sources (-want +got):\n%s", diff)
437438
}
438439

@@ -1864,7 +1865,7 @@ func TestInitializeConfigs(t *testing.T) {
18641865
ctx = util.WithInstrumentation(ctx, instrumentation)
18651866
t.Run("valid initialization", func(t *testing.T) {
18661867
sourceConfig1 := testutils.MockSourceConfig{Name: "my-source", Type: "mock-source"}
1867-
source1, _ := sourceConfig1.Initialize(ctx, nil)
1868+
source1, _ := sourceConfig1.Initialize(ctx, nil, false)
18681869
tools1 := testutils.NewMockTool("my-tool", "mock tool for offline config", "my-source", nil, false, false)
18691870
validCfg := server.ServerConfig{
18701871
Version: "0.0.0",

‎internal/sources/alloydbadmin/alloydbadmin.go‎

Lines changed: 46 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -63,48 +63,64 @@ func (r Config) SourceConfigType() string {
6363
return SourceType
6464
}
6565

66-
func (r Config) Initialize(ctx context.Context, tracer trace.Tracer) (sources.Source, error) {
67-
ua, err := util.UserAgentFromContext(ctx)
68-
if err != nil {
69-
return nil, fmt.Errorf("error in User Agent retrieval: %s", err)
70-
}
71-
72-
var client *http.Client
73-
if r.UseClientOAuth {
74-
client = &http.Client{
75-
Transport: util.NewUserAgentRoundTripper(ua, http.DefaultTransport),
76-
}
77-
} else {
78-
// Use Application Default Credentials
79-
creds, err := google.FindDefaultCredentials(ctx, alloydbrestapi.CloudPlatformScope)
80-
if err != nil {
81-
return nil, fmt.Errorf("failed to find default credentials: %w", err)
82-
}
83-
baseClient := oauth2.NewClient(ctx, creds.TokenSource)
84-
baseClient.Transport = util.NewUserAgentRoundTripper(ua, baseClient.Transport)
85-
client = baseClient
66+
func (r Config) Initialize(ctx context.Context, tracer trace.Tracer, deferConnect bool) (sources.Source, error) {
67+
s := r.newSource(ctx, tracer)
68+
if deferConnect {
69+
return s, nil
8670
}
87-
88-
service, err := alloydbrestapi.NewService(ctx, option.WithHTTPClient(client))
89-
if err != nil {
90-
return nil, fmt.Errorf("error creating new alloydb service: %w", err)
71+
if _, err := s.adminService(ctx); err != nil {
72+
return nil, err
9173
}
74+
return s, nil
75+
}
9276

93-
s := &Source{
77+
func (r Config) newSource(ctx context.Context, tracer trace.Tracer) *Source {
78+
return &Source{
9479
Config: r,
9580
BaseURL: "https://alloydb.googleapis.com",
96-
Service: service,
81+
tracer: tracer,
82+
conn: sources.NewConnectOnce[*alloydbrestapi.Service](ctx, r.Name, SourceType, tracer),
9783
}
98-
99-
return s, nil
10084
}
10185

10286
var _ sources.Source = &Source{}
10387

10488
type Source struct {
10589
Config
10690
BaseURL string
107-
Service *alloydbrestapi.Service
91+
tracer trace.Tracer
92+
conn *sources.ConnectOnce[*alloydbrestapi.Service]
93+
}
94+
95+
func (s *Source) adminService(ctx context.Context) (*alloydbrestapi.Service, error) {
96+
return s.conn.Do(ctx, func(ctx context.Context) (*alloydbrestapi.Service, error) {
97+
ua, err := util.UserAgentFromContext(ctx)
98+
if err != nil {
99+
return nil, fmt.Errorf("error in User Agent retrieval: %s", err)
100+
}
101+
102+
var client *http.Client
103+
if s.UseClientOAuth {
104+
client = &http.Client{
105+
Transport: util.NewUserAgentRoundTripper(ua, http.DefaultTransport),
106+
}
107+
} else {
108+
// Use Application Default Credentials
109+
creds, err := google.FindDefaultCredentials(ctx, alloydbrestapi.CloudPlatformScope)
110+
if err != nil {
111+
return nil, fmt.Errorf("failed to find default credentials: %w", err)
112+
}
113+
baseClient := oauth2.NewClient(ctx, creds.TokenSource)
114+
baseClient.Transport = util.NewUserAgentRoundTripper(ua, baseClient.Transport)
115+
client = baseClient
116+
}
117+
118+
service, err := alloydbrestapi.NewService(ctx, option.WithHTTPClient(client))
119+
if err != nil {
120+
return nil, fmt.Errorf("error creating new alloydb service: %w", err)
121+
}
122+
return service, nil
123+
})
108124
}
109125

110126
func (s *Source) IsReadOnly() bool {
@@ -133,7 +149,7 @@ func (s *Source) getService(ctx context.Context, accessToken string) (*alloydbre
133149
}
134150
return service, nil
135151
}
136-
return s.Service, nil
152+
return s.adminService(ctx)
137153
}
138154

139155
func (s *Source) UseClientAuthorization() bool {

‎internal/sources/alloydbpg/alloydb_pg.go‎

Lines changed: 47 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -68,35 +68,50 @@ func (r Config) SourceConfigType() string {
6868
return SourceType
6969
}
7070

71-
func (r Config) Initialize(ctx context.Context, tracer trace.Tracer) (sources.Source, error) {
72-
pool, err := initAlloyDBPgConnectionPool(ctx, tracer, r.Name, r.Project, r.Region, r.Cluster, r.Instance, r.IPType.String(), r.User, r.Password, r.Database, r.ReadOnly)
73-
if err != nil {
74-
return nil, fmt.Errorf("unable to create pool: %w", err)
71+
func (r Config) Initialize(ctx context.Context, tracer trace.Tracer, deferConnect bool) (sources.Source, error) {
72+
s := r.newSource(ctx, tracer)
73+
if deferConnect {
74+
return s, nil
7575
}
76-
77-
err = pool.Ping(ctx)
78-
if err != nil {
79-
pool.Close()
80-
if r.ReadOnly &&
81-
strings.Contains(err.Error(), "unrecognized configuration parameter") &&
82-
strings.Contains(err.Error(), "alloydb_session_read_only") {
83-
return nil, fmt.Errorf("failed to initialize AlloyDB source in read-only mode: 'alloydb_session_read_only' is not supported on this instance version. See documentation for details: https://mcp-toolbox.dev/integrations/alloydb/source/#reference: %w", err)
84-
}
85-
return nil, fmt.Errorf("unable to connect successfully: %w", err)
76+
if _, err := s.pool(ctx); err != nil {
77+
return nil, err
8678
}
79+
return s, nil
80+
}
8781

88-
s := &Source{
82+
func (r Config) newSource(ctx context.Context, tracer trace.Tracer) *Source {
83+
return &Source{
8984
Config: r,
90-
Pool: pool,
85+
conn: sources.NewConnectOnce[*pgxpool.Pool](ctx, r.Name, SourceType, tracer),
9186
}
92-
return s, nil
9387
}
9488

9589
var _ sources.Source = &Source{}
9690

9791
type Source struct {
9892
Config
99-
Pool *pgxpool.Pool
93+
conn *sources.ConnectOnce[*pgxpool.Pool]
94+
}
95+
96+
func (s *Source) pool(ctx context.Context) (*pgxpool.Pool, error) {
97+
return s.conn.Do(ctx, func(ctx context.Context) (*pgxpool.Pool, error) {
98+
pool, err := initAlloyDBPgConnectionPool(ctx, s.Project, s.Region, s.Cluster, s.Instance, s.IPType.String(), s.User, s.Password, s.Database, s.ReadOnly)
99+
if err != nil {
100+
return nil, fmt.Errorf("unable to create pool: %w", err)
101+
}
102+
103+
err = pool.Ping(ctx)
104+
if err != nil {
105+
pool.Close()
106+
if s.ReadOnly &&
107+
strings.Contains(err.Error(), "unrecognized configuration parameter") &&
108+
strings.Contains(err.Error(), "alloydb_session_read_only") {
109+
return nil, fmt.Errorf("failed to initialize AlloyDB source in read-only mode: 'alloydb_session_read_only' is not supported on this instance version. See documentation for details: https://mcp-toolbox.dev/integrations/alloydb/source/#reference: %w", err)
110+
}
111+
return nil, fmt.Errorf("unable to connect successfully: %w", err)
112+
}
113+
return pool, nil
114+
})
100115
}
101116

102117
func (s *Source) IsReadOnly() bool {
@@ -111,13 +126,24 @@ func (s *Source) ToConfig() sources.SourceConfig {
111126
return s.Config
112127
}
113128

129+
// PostgresPool reports the pool once connected; use PostgresPoolContext to guarantee one.
114130
func (s *Source) PostgresPool() *pgxpool.Pool {
115-
return s.Pool
131+
pool, _ := s.conn.Get()
132+
return pool
133+
}
134+
135+
// PostgresPoolContext returns the pool, connecting on first use.
136+
func (s *Source) PostgresPoolContext(ctx context.Context) (*pgxpool.Pool, error) {
137+
return s.pool(ctx)
116138
}
117139

118140
func (s *Source) RunSQL(ctx context.Context, statement string, params []any) (any, error) {
141+
pool, err := s.pool(ctx)
142+
if err != nil {
143+
return nil, err
144+
}
119145
statement = sqlcommenter.PrependComment(ctx, statement, SourceType, s.SQLCommenter)
120-
results, err := s.Pool.Query(ctx, statement, params...)
146+
results, err := pool.Query(ctx, statement, params...)
121147
if err != nil {
122148
return nil, fmt.Errorf("unable to execute query: %w", err)
123149
}
@@ -208,11 +234,7 @@ func getConnectionConfig(ctx context.Context, user, pass, dbname string, readOnl
208234
return dsn, useIAM, nil
209235
}
210236

211-
func initAlloyDBPgConnectionPool(ctx context.Context, tracer trace.Tracer, name, project, region, cluster, instance, ipType, user, pass, dbname string, readOnly bool) (*pgxpool.Pool, error) {
212-
//nolint:all // Reassigned ctx
213-
ctx, span := sources.InitConnectionSpan(ctx, tracer, SourceType, name)
214-
defer span.End()
215-
237+
func initAlloyDBPgConnectionPool(ctx context.Context, project, region, cluster, instance, ipType, user, pass, dbname string, readOnly bool) (*pgxpool.Pool, error) {
216238
dsn, useIAM, err := getConnectionConfig(ctx, user, pass, dbname, readOnly)
217239
if err != nil {
218240
return nil, fmt.Errorf("unable to get AlloyDB connection config: %w", err)

‎internal/sources/arcadedb/arcadedb.go‎

Lines changed: 36 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -77,30 +77,43 @@ func (r Config) SourceConfigType() string {
7777
return SourceType
7878
}
7979

80-
func (r Config) Initialize(ctx context.Context, tracer trace.Tracer) (sources.Source, error) {
81-
driver, err := initArcadeDBDriver(ctx, tracer, r.Uri, r.User, r.Password, r.Name)
82-
if err != nil {
83-
return nil, fmt.Errorf("unable to create driver: %w", err)
80+
func (r Config) Initialize(ctx context.Context, tracer trace.Tracer, deferConnect bool) (sources.Source, error) {
81+
s := r.newSource(ctx, tracer)
82+
if deferConnect {
83+
return s, nil
8484
}
85-
86-
err = driver.VerifyConnectivity(ctx)
87-
if err != nil {
88-
driver.Close(ctx)
89-
return nil, fmt.Errorf("unable to connect successfully: %w", err)
85+
if _, err := s.driver(ctx); err != nil {
86+
return nil, err
9087
}
88+
return s, nil
89+
}
9190

92-
s := &Source{
91+
func (r Config) newSource(ctx context.Context, tracer trace.Tracer) *Source {
92+
return &Source{
9393
Config: r,
94-
Driver: driver,
94+
conn: sources.NewConnectOnce[neo4j.Driver](ctx, r.Name, SourceType, tracer),
9595
}
96-
return s, nil
9796
}
9897

9998
var _ sources.Source = &Source{}
10099

101100
type Source struct {
102101
Config
103-
Driver neo4j.Driver
102+
conn *sources.ConnectOnce[neo4j.Driver]
103+
}
104+
105+
func (s *Source) driver(ctx context.Context) (neo4j.Driver, error) {
106+
return s.conn.Do(ctx, func(ctx context.Context) (neo4j.Driver, error) {
107+
driver, err := initArcadeDBDriver(ctx, s.Uri, s.User, s.Password)
108+
if err != nil {
109+
return nil, fmt.Errorf("unable to create driver: %w", err)
110+
}
111+
if err := driver.VerifyConnectivity(ctx); err != nil {
112+
driver.Close(ctx)
113+
return nil, fmt.Errorf("unable to connect successfully: %w", err)
114+
}
115+
return driver, nil
116+
})
104117
}
105118

106119
func (s *Source) IsReadOnly() bool {
@@ -115,8 +128,10 @@ func (s *Source) ToConfig() sources.SourceConfig {
115128
return s.Config
116129
}
117130

131+
// ArcadeDBDriver reports the driver once connected.
118132
func (s *Source) ArcadeDBDriver() neo4j.Driver {
119-
return s.Driver
133+
driver, _ := s.conn.Get()
134+
return driver
120135
}
121136

122137
func (s *Source) ArcadeDBDatabase() string {
@@ -137,8 +152,13 @@ func (s *Source) RunCypher(ctx context.Context, cypherStr string, params map[str
137152
cypherStr = "EXPLAIN " + cypherStr
138153
}
139154

155+
driver, err := s.driver(ctx)
156+
if err != nil {
157+
return nil, err
158+
}
159+
140160
config := neo4j.ExecuteQueryWithDatabase(s.ArcadeDBDatabase())
141-
results, err := neo4j.ExecuteQuery[*neo4j.EagerResult](ctx, s.ArcadeDBDriver(), cypherStr, params,
161+
results, err := neo4j.ExecuteQuery[*neo4j.EagerResult](ctx, driver, cypherStr, params,
142162
neo4j.EagerResultTransformer, config)
143163
if err != nil {
144164
return nil, fmt.Errorf("unable to execute query: %w", err)
@@ -315,10 +335,7 @@ func (s *Source) arcadeHTTPEndpointURL(endpoint string) (string, error) {
315335
url.PathEscape(s.ArcadeDBDatabase())), nil
316336
}
317337

318-
func initArcadeDBDriver(ctx context.Context, tracer trace.Tracer, uri, user, password, name string) (neo4j.Driver, error) {
319-
ctx, span := sources.InitConnectionSpan(ctx, tracer, SourceType, name)
320-
defer span.End()
321-
338+
func initArcadeDBDriver(ctx context.Context, uri, user, password string) (neo4j.Driver, error) {
322339
auth := neo4j.BasicAuth(user, password, "")
323340
userAgent, err := util.UserAgentFromContext(ctx)
324341
if err != nil {

0 commit comments

Comments
 (0)