feat: synchronise provider changes with WatchAIProviders (#27091)

## Why

PR #26797 was accidentally merged into the stale `graphite-base/26797`
branch instead of `main` (Graphite picked the wrong base), so its
changes never landed on `main`. This PR re-lands that work as a clean
cherry-pick onto the current `main`.

## What

Adds a `WatchAIProviders` streaming RPC to the `ProviderConfigurator`
service so a running standalone AI Gateway refetches its provider set
when the provider configuration changes. The server subscribes to
`AIProvidersChangedChannel` (published by the provider CRUD endpoints)
and forwards each event as a payload-free signal, plus one signal on
subscribe; the gateway calls `GetAIProviders` on each signal to rebuild
its pool. The aibridged API is bumped to v1.2.

Env-seeded providers don't need a signal: seeding finishes before coderd
serves the gateway connection, so the gateway's initial fetch already
reflects the seeded set.

## For reviewers

The change is split into two commits to make review easy:

1. **`feat: synchronise provider changes with WatchAIProviders`** is a
faithful cherry-pick of #26797, identical to the originally reviewed PR.
It is committed without pre-commit hooks because it does not build
against current `main` on its own.
2. **`fix: resolve cherry-pick conflicts against main`** contains only
the deltas needed to re-land on current `main`, and passes the full
pre-commit suite:
- `coderd/aibridged/proto/aibridged.pb.go` regenerated via the proto
make target (the cherry-picked copy was generated against the older
proto).
- `enterprise/cli/aigatewaystart.go` import block unioned; `main` added
`os` and `strings` while the PR added `sync`.
- Three `aibridgedserver.NewServer` test call sites that landed on
`main` after the original branch diverged now pass the new `pubsub`
argument.

Refs https://linear.app/codercom/issue/AIGOV-465

*This PR was produced by opencode (agent) using the
`anthropic/claude-opus-4-8` model, under human direction and review.*
This commit is contained in:
Danny Kopping
2026-07-08 15:32:17 +02:00
committed by GitHub
parent 48f07e6e13
commit affb359d13
16 changed files with 952 additions and 191 deletions
+57 -1
View File
@@ -26,9 +26,11 @@ import (
"github.com/coder/coder/v2/coderd/database/db2sdk"
"github.com/coder/coder/v2/coderd/database/dbauthz"
"github.com/coder/coder/v2/coderd/database/dbtime"
"github.com/coder/coder/v2/coderd/database/pubsub"
"github.com/coder/coder/v2/coderd/externalauth"
"github.com/coder/coder/v2/coderd/httpmw"
codermcp "github.com/coder/coder/v2/coderd/mcp"
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
"github.com/coder/coder/v2/coderd/util/ptr"
"github.com/coder/coder/v2/codersdk"
)
@@ -101,6 +103,7 @@ type Server struct {
// long-running operations.
lifecycleCtx context.Context
store store
pubsub pubsub.Pubsub
logger slog.Logger
externalAuthConfigs map[string]*externalauth.Config
@@ -112,7 +115,7 @@ type Server struct {
budgetPolicy codersdk.AIBudgetPolicy
}
func NewServer(lifecycleCtx context.Context, store store, logger slog.Logger, accessURL string,
func NewServer(lifecycleCtx context.Context, store store, ps pubsub.Pubsub, logger slog.Logger, accessURL string,
bridgeCfg codersdk.AIBridgeConfig, externalAuthConfigs []*externalauth.Config, experiments codersdk.Experiments,
aiSeatTracker aiseats.SeatTracker,
) (*Server, error) {
@@ -129,6 +132,7 @@ func NewServer(lifecycleCtx context.Context, store store, logger slog.Logger, ac
srv := &Server{
lifecycleCtx: lifecycleCtx,
store: store,
pubsub: ps,
logger: logger,
externalAuthConfigs: eac,
structuredLogging: bridgeCfg.StructuredLogging.Value(),
@@ -905,6 +909,58 @@ func (s *Server) GetAIProviders(ctx context.Context, _ *proto.GetAIProvidersRequ
return &proto.GetAIProvidersResponse{Providers: providers}, nil
}
// WatchAIProviders streams a payload-free change signal on each
// AIProvidersChangedChannel event, plus one immediately on subscribe. Pubsub
// drop errors produce a signal rather than failing the stream. Blocks until the
// stream context or the server lifecycle is canceled.
func (s *Server) WatchAIProviders(_ *proto.WatchAIProvidersRequest, stream proto.DRPCProviderConfigurator_WatchAIProvidersStream) error {
if s.pubsub == nil {
return xerrors.New("pubsub not configured")
}
ctx, cancel := context.WithCancel(stream.Context())
defer cancel()
// Cancel when the server lifecycle ends, not just when the stream closes.
stop := context.AfterFunc(s.lifecycleCtx, cancel)
defer stop()
// Buffered to one so a burst of events collapses into a single pending
// signal.
signals := make(chan struct{}, 1)
notify := func() {
select {
case signals <- struct{}{}:
default:
}
}
// Every event signals, including dropped-message errors.
unsubscribe, err := s.pubsub.SubscribeWithErr(coderdpubsub.AIProvidersChangedChannel, func(cbCtx context.Context, _ []byte, err error) {
if err != nil {
s.logger.Warn(cbCtx, "ai providers changed event delivered with error", slog.Error(err))
}
notify()
})
if err != nil {
return xerrors.Errorf("subscribe to %s: %w", coderdpubsub.AIProvidersChangedChannel, err)
}
defer unsubscribe()
// Initial signal on subscribe.
notify()
for {
select {
case <-ctx.Done():
return nil
case <-signals:
if err := stream.Send(&proto.WatchAIProvidersResponse{}); err != nil {
return xerrors.Errorf("send ai providers change signal: %w", err)
}
}
}
}
// Deprecated: Injected MCP in AI Bridge is deprecated and will be removed in a future release.
func getCoderMCPServerConfig(experiments codersdk.Experiments, accessURL string) (*proto.MCPServerConfig, error) {
// Both the MCP & OAuth2 experiments are currently required in order to use our
+176 -13
View File
@@ -19,10 +19,12 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"golang.org/x/xerrors"
protobufproto "google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/anypb"
"google.golang.org/protobuf/types/known/structpb"
"google.golang.org/protobuf/types/known/timestamppb"
"storj.io/drpc"
"cdr.dev/slog/v3"
"cdr.dev/slog/v3/sloggers/slogjson"
@@ -40,8 +42,10 @@ import (
"github.com/coder/coder/v2/coderd/database/dbmock"
"github.com/coder/coder/v2/coderd/database/dbtestutil"
"github.com/coder/coder/v2/coderd/database/dbtime"
"github.com/coder/coder/v2/coderd/database/pubsub"
"github.com/coder/coder/v2/coderd/externalauth"
codermcp "github.com/coder/coder/v2/coderd/mcp"
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
"github.com/coder/coder/v2/coderd/rbac"
"github.com/coder/coder/v2/coderd/util/ptr"
"github.com/coder/coder/v2/codersdk"
@@ -205,7 +209,7 @@ func TestAuthorization(t *testing.T) {
tc.mocksFn(db, apiKey, user)
}
srv, err := aibridgedserver.NewServer(t.Context(), db, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
srv, err := aibridgedserver.NewServer(t.Context(), db, nil, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
require.NotNil(t, srv)
@@ -367,7 +371,7 @@ func TestAuthorization_Delegated(t *testing.T) {
tc.mocksFn(db, apiKey, user)
}
srv, err := aibridgedserver.NewServer(t.Context(), db, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
srv, err := aibridgedserver.NewServer(t.Context(), db, nil, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
require.NotNil(t, srv)
@@ -563,7 +567,7 @@ func TestIsBudgetExceeded(t *testing.T) {
wantResp = tc.setupMocks(db, userID)
}
srv, err := aibridgedserver.NewServer(t.Context(), db, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
srv, err := aibridgedserver.NewServer(t.Context(), db, nil, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
req := &proto.IsBudgetExceededRequest{UserId: userIDStr}
@@ -616,7 +620,7 @@ func TestIsBudgetExceeded_Enforcement(t *testing.T) {
})
require.NoError(t, err)
srv, err := aibridgedserver.NewServer(ctx, authzDB, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
srv, err := aibridgedserver.NewServer(ctx, authzDB, nil, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
return ctx, rawDB, srv, user, group
@@ -771,7 +775,7 @@ func TestGetMCPServerConfigs(t *testing.T) {
logger := testutil.Logger(t)
accessURL := "https://my-cool-deployment.com"
srv, err := aibridgedserver.NewServer(t.Context(), db, logger, accessURL, codersdk.AIBridgeConfig{
srv, err := aibridgedserver.NewServer(t.Context(), db, nil, logger, accessURL, codersdk.AIBridgeConfig{
InjectCoderMCPTools: serpent.Bool(!tc.disableCoderMCPInjection),
}, tc.externalAuthConfigs, tc.experiments, agplaiseats.Noop{})
require.NoError(t, err)
@@ -811,7 +815,7 @@ func TestGetMCPServerAccessTokensBatch(t *testing.T) {
logger := testutil.Logger(t)
// Given: 2 external auth configured with MCP and 1 without.
srv, err := aibridgedserver.NewServer(t.Context(), db, logger, "/", codersdk.AIBridgeConfig{}, []*externalauth.Config{
srv, err := aibridgedserver.NewServer(t.Context(), db, nil, logger, "/", codersdk.AIBridgeConfig{}, []*externalauth.Config{
{
ID: "1",
MCPURL: "1.com/mcp",
@@ -2032,7 +2036,7 @@ func TestRecordTokenUsageAuthorized(t *testing.T) {
now := time.Date(2026, 6, 25, 14, 30, 0, 0, time.UTC)
// The server runs every store call as subjectAibridged via the authzDB.
srv, err := aibridgedserver.NewServer(ctx, authzDB, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
srv, err := aibridgedserver.NewServer(ctx, authzDB, nil, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
_, err = srv.RecordTokenUsage(ctx, &proto.RecordTokenUsageRequest{
@@ -2407,7 +2411,7 @@ func testRecordMethod[Req any, Resp any](
}
ctx := testutil.Context(t, testutil.WaitLong)
srv, err := aibridgedserver.NewServer(ctx, db, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
srv, err := aibridgedserver.NewServer(ctx, db, nil, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
resp, err := callMethod(srv, ctx, tc.request)
@@ -2727,7 +2731,7 @@ func TestStructuredLogging(t *testing.T) {
tc.setupMocks(db, interceptionID)
ctx := testutil.Context(t, testutil.WaitLong)
srv, err := aibridgedserver.NewServer(ctx, db, logger, "/", codersdk.AIBridgeConfig{
srv, err := aibridgedserver.NewServer(ctx, db, nil, logger, "/", codersdk.AIBridgeConfig{
StructuredLogging: serpent.Bool(tc.structuredLogging),
}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
@@ -2771,7 +2775,7 @@ func TestInferredThreadsByToolCalls(t *testing.T) {
user := dbgen.User(t, db, database.User{})
srv, err := aibridgedserver.NewServer(ctx, db, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
srv, err := aibridgedserver.NewServer(ctx, db, nil, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
aID := uuid.New()
@@ -2867,7 +2871,7 @@ func TestRecordToolUsageProviderItemID(t *testing.T) {
user := dbgen.User(t, db, database.User{})
srv, err := aibridgedserver.NewServer(ctx, db, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
srv, err := aibridgedserver.NewServer(ctx, db, nil, logger, "/", codersdk.AIBridgeConfig{}, nil, requiredExperiments, agplaiseats.Noop{})
require.NoError(t, err)
intcID := uuid.New()
@@ -2997,7 +3001,7 @@ func TestGetAIProviders(t *testing.T) {
Settings: sql.NullString{String: "{not valid json", Valid: true},
})
srv, err := aibridgedserver.NewServer(ctx, db, logger, "/", codersdk.AIBridgeConfig{}, nil, nil, agplaiseats.Noop{})
srv, err := aibridgedserver.NewServer(ctx, db, nil, logger, "/", codersdk.AIBridgeConfig{}, nil, nil, agplaiseats.Noop{})
require.NoError(t, err)
resp, err := srv.GetAIProviders(ctx, &proto.GetAIProvidersRequest{})
@@ -3060,7 +3064,7 @@ func TestGetAIProvidersBlocksOnSeedLock(t *testing.T) {
BaseUrl: "https://api.openai.com/",
}, "sk-openai")
srv, err := aibridgedserver.NewServer(ctx, db, logger, "/", codersdk.AIBridgeConfig{}, nil, nil, agplaiseats.Noop{})
srv, err := aibridgedserver.NewServer(ctx, db, nil, logger, "/", codersdk.AIBridgeConfig{}, nil, nil, agplaiseats.Noop{})
require.NoError(t, err)
// Simulate an in-flight env seed holding the advisory lock until released.
@@ -3128,3 +3132,162 @@ func TestGetAIProvidersBlocksOnSeedLock(t *testing.T) {
assert.Equal(t, "openai", resp.GetProviders()[0].GetName())
assert.Equal(t, []string{"sk-openai"}, resp.GetProviders()[0].GetKeys())
}
// TestWatchAIProviders asserts that the WatchAIProviders handler emits an
// initial signal on subscribe, one signal per AIProvidersChangedChannel publish,
// and returns cleanly when the stream context is canceled.
func TestWatchAIProviders(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitLong)
logger := slogtest.Make(t, nil)
// In-memory pubsub delivers Publish synchronously for deterministic signals.
ps := pubsub.NewInMemory()
srv, err := aibridgedserver.NewServer(ctx, db, ps, logger, "/", codersdk.AIBridgeConfig{}, nil, nil, agplaiseats.Noop{})
require.NoError(t, err)
streamCtx, streamCancel := context.WithCancel(ctx)
defer streamCancel()
stream := &fakeWatchProvidersStream{ctx: streamCtx, sent: make(chan struct{}, 16)}
watchErr := make(chan error, 1)
go func() {
watchErr <- srv.WatchAIProviders(&proto.WatchAIProvidersRequest{}, stream)
}()
// The handler sends an initial signal immediately on subscribe. Draining it
// before publishing guarantees the next publish is not coalesced into the
// initial signal.
testutil.TryReceive(ctx, t, stream.sent)
require.NoError(t, ps.Publish(coderdpubsub.AIProvidersChangedChannel, nil))
testutil.TryReceive(ctx, t, stream.sent)
require.NoError(t, ps.Publish(coderdpubsub.AIProvidersChangedChannel, nil))
testutil.TryReceive(ctx, t, stream.sent)
streamCancel()
require.NoError(t, testutil.TryReceive(ctx, t, watchErr))
}
// TestWatchAIProvidersSignalsOnDeliveryError asserts that a dropped-message
// delivery error is forwarded as a change signal rather than failing the
// stream, so the gateway reconverges after a pubsub drop.
func TestWatchAIProvidersSignalsOnDeliveryError(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitLong)
logger := slogtest.Make(t, nil)
ps := &captureListenerPubsub{listenerC: make(chan pubsub.ListenerWithErr, 1)}
srv, err := aibridgedserver.NewServer(ctx, db, ps, logger, "/", codersdk.AIBridgeConfig{}, nil, nil, agplaiseats.Noop{})
require.NoError(t, err)
streamCtx, streamCancel := context.WithCancel(ctx)
defer streamCancel()
stream := &fakeWatchProvidersStream{ctx: streamCtx, sent: make(chan struct{}, 16)}
watchErr := make(chan error, 1)
go func() {
watchErr <- srv.WatchAIProviders(&proto.WatchAIProvidersRequest{}, stream)
}()
// Capture the registered listener and drain the initial subscribe signal so
// the delivery-error signal that follows is not coalesced into it.
listener := testutil.TryReceive(ctx, t, ps.listenerC)
testutil.TryReceive(ctx, t, stream.sent)
// A delivery error must still produce a signal, exercising the pubsub-error
// branch of the handler.
listener(ctx, nil, pubsub.ErrDroppedMessages)
testutil.TryReceive(ctx, t, stream.sent)
streamCancel()
require.NoError(t, testutil.TryReceive(ctx, t, watchErr))
}
// TestWatchAIProvidersStopsOnLifecycleCancel asserts the handler returns when
// the server lifecycle context is canceled even though the stream context
// remains open, so a stream that outlives the server does not leak a goroutine
// on shutdown.
func TestWatchAIProvidersStopsOnLifecycleCancel(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitLong)
logger := slogtest.Make(t, nil)
ps := pubsub.NewInMemory()
// The lifecycle context is independent of the stream context so it can be
// canceled while the stream stays open.
lifecycleCtx, lifecycleCancel := context.WithCancel(ctx)
defer lifecycleCancel()
srv, err := aibridgedserver.NewServer(lifecycleCtx, db, ps, logger, "/", codersdk.AIBridgeConfig{}, nil, nil, agplaiseats.Noop{})
require.NoError(t, err)
streamCtx, streamCancel := context.WithCancel(ctx)
defer streamCancel()
stream := &fakeWatchProvidersStream{ctx: streamCtx, sent: make(chan struct{}, 16)}
watchErr := make(chan error, 1)
go func() {
watchErr <- srv.WatchAIProviders(&proto.WatchAIProvidersRequest{}, stream)
}()
// Drain the initial subscribe signal to confirm the handler is running
// before the lifecycle is canceled.
testutil.TryReceive(ctx, t, stream.sent)
// Canceling only the lifecycle context must stop the handler even though
// the stream context is still open.
lifecycleCancel()
require.NoError(t, testutil.TryReceive(ctx, t, watchErr))
}
var _ pubsub.Pubsub = (*captureListenerPubsub)(nil)
// captureListenerPubsub captures the ListenerWithErr registered via
// SubscribeWithErr so a test can drive delivery (including errors) directly.
type captureListenerPubsub struct {
listenerC chan pubsub.ListenerWithErr
}
func (*captureListenerPubsub) Subscribe(string, pubsub.Listener) (func(), error) {
return nil, xerrors.New("Subscribe not implemented")
}
func (p *captureListenerPubsub) SubscribeWithErr(_ string, listener pubsub.ListenerWithErr) (func(), error) {
p.listenerC <- listener
return func() {}, nil
}
func (*captureListenerPubsub) Publish(string, []byte) error {
return xerrors.New("Publish not implemented")
}
func (*captureListenerPubsub) Close() error { return nil }
// fakeWatchProvidersStream is a minimal proto.DRPCProviderConfigurator_WatchAIProvidersStream
// that records Send calls on a channel.
type fakeWatchProvidersStream struct {
ctx context.Context
sent chan struct{}
}
func (s *fakeWatchProvidersStream) Send(*proto.WatchAIProvidersResponse) error {
select {
case s.sent <- struct{}{}:
return nil
case <-s.ctx.Done():
return s.ctx.Err()
}
}
func (s *fakeWatchProvidersStream) Context() context.Context { return s.ctx }
func (*fakeWatchProvidersStream) MsgSend(drpc.Message, drpc.Encoding) error { return nil }
func (*fakeWatchProvidersStream) MsgRecv(drpc.Message, drpc.Encoding) error { return nil }
func (*fakeWatchProvidersStream) CloseSend() error { return nil }
func (*fakeWatchProvidersStream) Close() error { return nil }