mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: remove unbound Client() method from aibridged.Server (#27845)
Adds client context to `Client()` method in `aibridged.Server`, effectivly renaming `ClientContext()` method as `Client()`. Similarly `aibridged.ClientFuncWithContext` became `aibridged.ClientFunc`. `aibridged.Server.Client()` acquired a DRPC client with `context.Background()`, callers in theory could wait indefinitely for the daemon to connect to coderd. Every call site already had a context except the recorder callback. `aibridge.NewRecorder` takes a `func(context.Context) (Recorder, error)` and acquires against the record call's context.
This commit is contained in:
@@ -167,11 +167,9 @@ func (s *Server) Err() error {
|
||||
return s.lifecycleCtx.Err()
|
||||
}
|
||||
|
||||
func (s *Server) Client() (DRPCClient, error) {
|
||||
return s.ClientContext(context.Background())
|
||||
}
|
||||
|
||||
func (s *Server) ClientContext(ctx context.Context) (DRPCClient, error) {
|
||||
// Client acquires a [DRPCClient], blocking until the daemon is connected to
|
||||
// coderd, the server lifecycle ends, or ctx is canceled.
|
||||
func (s *Server) Client(ctx context.Context) (DRPCClient, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
|
||||
@@ -123,7 +123,7 @@ func TestClient_TransientDialErrorRetries(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = srv.Shutdown(context.Background()) })
|
||||
|
||||
_, err = srv.ClientContext(testutil.Context(t, testutil.WaitShort))
|
||||
_, err = srv.Client(testutil.Context(t, testutil.WaitShort))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int32(2), calls.Load())
|
||||
}
|
||||
|
||||
@@ -10,12 +10,10 @@ import (
|
||||
|
||||
type Dialer func(ctx context.Context) (DRPCClient, error)
|
||||
|
||||
type ClientFunc func() (DRPCClient, error)
|
||||
|
||||
// ClientFuncWithContext acquires a DRPCClient, honoring the passed context so a
|
||||
// blocking acquisition (e.g. waiting for the daemon to connect to coderd)
|
||||
// unblocks when the context is canceled. Server.ClientContext satisfies it.
|
||||
type ClientFuncWithContext func(context.Context) (DRPCClient, error)
|
||||
// ClientFunc acquires a DRPCClient, honoring the passed context so a blocking
|
||||
// acquisition (e.g. waiting for the daemon to connect to coderd) unblocks when
|
||||
// the context is canceled. Server.Client satisfies it.
|
||||
type ClientFunc func(context.Context) (DRPCClient, error)
|
||||
|
||||
// DRPCClient is the union of various service interfaces the client must support.
|
||||
type DRPCClient interface {
|
||||
|
||||
@@ -113,7 +113,7 @@ func (s *Server) ServeHTTP(rw http.ResponseWriter, r *http.Request) {
|
||||
r.Header.Del("X-Api-Key")
|
||||
}
|
||||
|
||||
client, err := s.ClientContext(ctx)
|
||||
client, err := s.Client(ctx)
|
||||
if err != nil {
|
||||
logger.Warn(ctx, "failed to connect to coderd", slog.Error(err))
|
||||
http.Error(rw, ErrConnect.Error(), http.StatusServiceUnavailable)
|
||||
|
||||
@@ -60,7 +60,7 @@ func (m *MCPProxyFactory) Build(ctx context.Context, req Request, tracer trace.T
|
||||
}
|
||||
|
||||
func (m *MCPProxyFactory) retrieveMCPServerConfigs(ctx context.Context, req Request) (map[string]mcp.ServerProxier, error) {
|
||||
client, err := m.clientFn()
|
||||
client, err := m.clientFn(ctx)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("acquire client: %w", err)
|
||||
}
|
||||
|
||||
@@ -228,8 +228,10 @@ func (p *CachedBridgePool) Acquire(ctx context.Context, req Request, clientFn Cl
|
||||
|
||||
span.AddEvent("cache_miss")
|
||||
providerVersion := p.providerVersion.Load()
|
||||
recorder := aibridge.NewRecorder(p.logger.Named("recorder"), p.tracer, func() (aibridge.Recorder, error) {
|
||||
client, err := clientFn()
|
||||
recorder := aibridge.NewRecorder(p.logger.Named("recorder"), p.tracer, func(clientCtx context.Context) (aibridge.Recorder, error) {
|
||||
// The recorder outlives this Acquire call, so the client is acquired
|
||||
// against the context of the record call being served.
|
||||
client, err := clientFn(clientCtx)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("acquire client: %w", err)
|
||||
}
|
||||
|
||||
@@ -47,7 +47,7 @@ func TestPool(t *testing.T) {
|
||||
t.Cleanup(func() { pool.Shutdown(context.Background()) })
|
||||
|
||||
id, id2, apiKeyID1, apiKeyID2 := uuid.New(), uuid.New(), uuid.New(), uuid.New()
|
||||
clientFn := func() (aibridged.DRPCClient, error) {
|
||||
clientFn := func(context.Context) (aibridged.DRPCClient, error) {
|
||||
return client, nil
|
||||
}
|
||||
|
||||
@@ -149,7 +149,7 @@ func TestPoolReplaceProvidersClearsCacheAndUsesNewProviders(t *testing.T) {
|
||||
InitiatorID: uuid.New(),
|
||||
APIKeyID: uuid.New().String(),
|
||||
}
|
||||
clientFn := func() (aibridged.DRPCClient, error) {
|
||||
clientFn := func(context.Context) (aibridged.DRPCClient, error) {
|
||||
return client, nil
|
||||
}
|
||||
|
||||
@@ -195,7 +195,7 @@ func TestPoolReplaceProvidersDoesNotJoinStaleSingleflight(t *testing.T) {
|
||||
InitiatorID: uuid.New(),
|
||||
APIKeyID: uuid.New().String(),
|
||||
}
|
||||
clientFn := func() (aibridged.DRPCClient, error) {
|
||||
clientFn := func(context.Context) (aibridged.DRPCClient, error) {
|
||||
return client, nil
|
||||
}
|
||||
|
||||
@@ -266,7 +266,7 @@ func TestPoolReplaceProvidersAfterShutdownIsNoop(t *testing.T) {
|
||||
SessionKey: "key",
|
||||
InitiatorID: uuid.New(),
|
||||
APIKeyID: uuid.New().String(),
|
||||
}, func() (aibridged.DRPCClient, error) {
|
||||
}, func(context.Context) (aibridged.DRPCClient, error) {
|
||||
return nil, context.Canceled
|
||||
}, newMockMCPFactory(nil))
|
||||
require.ErrorContains(t, err, "pool shutting down")
|
||||
@@ -294,7 +294,7 @@ func TestPool_Expiry(t *testing.T) {
|
||||
InitiatorID: uuid.New(),
|
||||
APIKeyID: uuid.New().String(),
|
||||
}
|
||||
clientFn := func() (aibridged.DRPCClient, error) {
|
||||
clientFn := func(context.Context) (aibridged.DRPCClient, error) {
|
||||
return client, nil
|
||||
}
|
||||
|
||||
@@ -384,7 +384,7 @@ func TestPoolShutdownReplaceProviders(t *testing.T) {
|
||||
}, logger, nil, testTracer)
|
||||
require.NoError(t, err)
|
||||
|
||||
clientFn := func() (aibridged.DRPCClient, error) { return client, nil }
|
||||
clientFn := func(context.Context) (aibridged.DRPCClient, error) { return client, nil }
|
||||
|
||||
// Populate the cache so ReplaceProviders' Clear has an entry to evict.
|
||||
_, err = pool.Acquire(ctx, aibridged.Request{
|
||||
|
||||
@@ -70,12 +70,12 @@ func SubscribeProviderReload(
|
||||
// before serving.
|
||||
//
|
||||
// It runs until ctx is canceled, then returns ctx.Err(). clientFn receives ctx,
|
||||
// so a client acquisition that blocks (e.g. Server.ClientContext waiting for
|
||||
// so a client acquisition that blocks (e.g. Server.Client waiting for
|
||||
// the daemon to connect to coderd) unblocks when ctx is canceled, leaving no
|
||||
// goroutine behind.
|
||||
func WatchProviderReload(
|
||||
ctx context.Context,
|
||||
clientFn ClientFuncWithContext,
|
||||
clientFn ClientFunc,
|
||||
reloader ProviderReloader,
|
||||
logger slog.Logger,
|
||||
) error {
|
||||
@@ -109,7 +109,7 @@ func WatchProviderReload(
|
||||
// watchProviderReloadOnce opens a single WatchAIProviders stream and reloads on
|
||||
// each signal until the stream fails. received reports whether at least one
|
||||
// signal was received before the error.
|
||||
func watchProviderReloadOnce(ctx context.Context, clientFn ClientFuncWithContext, reloader ProviderReloader, logger slog.Logger) (received bool, err error) {
|
||||
func watchProviderReloadOnce(ctx context.Context, clientFn ClientFunc, reloader ProviderReloader, logger slog.Logger) (received bool, err error) {
|
||||
// clientFn blocks until the daemon connects to coderd or ctx is canceled.
|
||||
c, err := clientFn(ctx)
|
||||
if err != nil {
|
||||
|
||||
@@ -186,7 +186,7 @@ func TestWatchProviderReloadCancelUnblocksClient(t *testing.T) {
|
||||
logger := slogtest.Make(t, nil)
|
||||
|
||||
// clientFn blocks until its context is canceled, modeling
|
||||
// Server.ClientContext waiting for the daemon to connect to coderd. Only
|
||||
// Server.Client waiting for the daemon to connect to coderd. Only
|
||||
// watchCancel is exercised (no stream activity, no daemon close), so the
|
||||
// loop can return only if clientFn honors the context it receives.
|
||||
var once sync.Once
|
||||
|
||||
@@ -78,7 +78,7 @@ func StartTestAIBridgeDaemonWithPubsub(
|
||||
|
||||
// The reloader fetches providers from coderd over srv's DRPC client; the
|
||||
// subscription drives an initial load and refreshes on change events.
|
||||
reloader := cli.NewPoolRPCReloader(pool, srv.ClientContext, cfg, logger.Named("reloader"), nil, metrics)
|
||||
reloader := cli.NewPoolRPCReloader(pool, srv.Client, cfg, logger.Named("reloader"), nil, metrics)
|
||||
unsubscribe, err := aibridged.SubscribeProviderReload(ctx, ps, reloader, logger.Named("subscriber"))
|
||||
if err != nil {
|
||||
t.Fatalf("subscribe provider reload: %v", err)
|
||||
|
||||
Reference in New Issue
Block a user