feat: fetch providers over DRPC (#26650)

Closes [AIGOV-455](https://linear.app/codercom/issue/AIGOV-455/extend-drpc-with-buildproviders).

## Why

The AI Gateway (`aibridged`) is being split into a standalone process that must not touch the database. `coderd` stays the source of truth and seeds the `ai_providers` / `ai_provider_keys` tables from the environment. This PR adds a DRPC call so the gateway fetches provider config from `coderd` instead of reading the DB, for both the embedded and standalone daemons.

## What

- **Proto:** new `ProviderConfigurator` service with a unary `GetAIProviders` RPC, plus `AIProvider` / `AIProviderBedrock` messages. `CurrentMinor` bumped to 1 (additive).
- **Server (`coderd/aibridgedserver`):** `GetAIProviders` runs a read-only `InTx` under `LockIDAIProvidersEnvSeed` so it never returns a mid-seed snapshot, reads providers (incl. disabled) plus keys for enabled ones, and maps to proto under `dbauthz.AsAIBridged`. Unmappable rows are skipped and logged; plaintext keys and Bedrock secrets are never logged.
- **Client:** `DRPCProviderConfiguratorClient` wired into the client union, `dialer.go`, and `CreateInMemoryAIBridgeServer`.
- **cli:** `BuildProvidersFromProto` maps the response through the existing DB-neutral `buildProvider`. A shared `poolRPCReloader` does the fetch/build/replace for both daemons: the embedded daemon reloads on every `ai_providers` change and fails startup if it cannot subscribe; the standalone gateway drives the same reloader once at startup, retrying until success and staying interruptible.
- **Dead code removed:** `BuildProvidersFromConfig`, `ProvidersFromConfig`, `AIProviderFromConfig`, and the DB-read `BuildProviders` path.
This commit is contained in:
Danny Kopping
2026-06-29 13:34:58 +02:00
committed by GitHub
parent 74b8f10d4e
commit ce94d42e19
22 changed files with 1288 additions and 408 deletions
+137
View File
@@ -21,11 +21,13 @@ import (
"github.com/coder/coder/v2/coderd/aiseats"
"github.com/coder/coder/v2/coderd/apikey"
"github.com/coder/coder/v2/coderd/database"
"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/externalauth"
"github.com/coder/coder/v2/coderd/httpmw"
codermcp "github.com/coder/coder/v2/coderd/mcp"
"github.com/coder/coder/v2/coderd/util/ptr"
"github.com/coder/coder/v2/codersdk"
)
@@ -80,6 +82,13 @@ type store interface {
// Authorizer-related queries.
GetAPIKeyByID(ctx context.Context, id string) (database.APIKey, error)
GetUserByID(ctx context.Context, id uuid.UUID) (database.User, error)
// ProviderConfigurator-related queries. InTx wraps the provider and key
// reads in a single read-only transaction; AcquireLock serializes against
// any in-flight env seed holding LockIDAIProvidersEnvSeed.
InTx(func(database.Store) error, *database.TxOptions) error
GetAIProviders(ctx context.Context, arg database.GetAIProvidersParams) ([]database.AIProvider, error)
GetAIProviderKeysByProviderIDs(ctx context.Context, providerIDs []uuid.UUID) ([]database.AIProviderKey, error)
}
type Server struct {
@@ -682,6 +691,93 @@ func (s *Server) IsAuthorized(ctx context.Context, in *proto.IsAuthorizedRequest
}, nil
}
// GetAIProviders returns the full AI provider set (enabled and disabled) from
// the database, which is the single source of truth seeded from coderd's
// environment. Embedded and standalone AI Gateway daemons call this over DRPC
// to build their provider pool instead of reading the database directly.
//
// The handler reads under a read-only transaction that first acquires
// LockIDAIProvidersEnvSeed, so it blocks until any in-flight env seed commits
// or rolls back. This guarantees the response is never a partial, mid-seed
// snapshot.
//
// Keys are populated only for enabled providers; disabled providers never call
// upstream, so their secrets are withheld.
//
// SECURITY: the response carries plaintext API keys and Bedrock credentials.
// Do not log the response struct.
func (s *Server) GetAIProviders(ctx context.Context, _ *proto.GetAIProvidersRequest) (*proto.GetAIProvidersResponse, error) {
//nolint:gocritic // AIBridged has a minimal permission set scoped to AI Bridge queries.
ctx = dbauthz.AsAIBridged(ctx)
var (
rows []database.AIProvider
keysByProvider map[uuid.UUID][]database.AIProviderKey
)
// Wrap both reads in a read-only transaction so the provider list and the
// key list are consistent with each other, and so the seed lock is held
// for the duration of the reads.
err := s.store.InTx(func(tx database.Store) error {
// Block on any in-flight seed transaction holding the advisory lock so
// the response reflects a fully-seeded snapshot.
if err := tx.AcquireLock(ctx, database.LockIDAIProvidersEnvSeed); err != nil {
return xerrors.Errorf("acquire ai providers env seed lock: %w", err)
}
var err error
rows, err = tx.GetAIProviders(ctx, database.GetAIProvidersParams{IncludeDisabled: true})
if err != nil {
return xerrors.Errorf("get ai providers: %w", err)
}
// Load keys only for enabled providers to avoid materializing secrets
// for disabled rows.
ids := make([]uuid.UUID, 0, len(rows))
for _, row := range rows {
if !row.Enabled {
continue
}
ids = append(ids, row.ID)
}
keysByProvider = make(map[uuid.UUID][]database.AIProviderKey, len(ids))
if len(ids) == 0 {
return nil
}
keyRows, err := tx.GetAIProviderKeysByProviderIDs(ctx, ids)
if err != nil {
return xerrors.Errorf("get ai provider keys: %w", err)
}
for _, k := range keyRows {
keysByProvider[k.ProviderID] = append(keysByProvider[k.ProviderID], k)
}
return nil
}, &database.TxOptions{ReadOnly: true, TxIdentifier: "get_ai_providers"})
if err != nil {
return nil, err
}
providers := make([]*proto.AIProvider, 0, len(rows))
for _, row := range rows {
p, err := aiProviderToProto(row, keysByProvider[row.ID])
if err != nil {
// Skip the offending row rather than failing the whole fetch:
// one row with a corrupt settings blob must not break provider
// configuration for every gateway, which would otherwise loop
// forever on the empty pool.
s.logger.Error(ctx, "skipping ai provider with invalid settings; it will be absent from the gateway pool",
slog.F("provider_id", row.ID),
slog.F("provider_name", row.Name),
slog.F("provider_type", string(row.Type)),
slog.Error(err),
)
continue
}
providers = append(providers, p)
}
return &proto.GetAIProvidersResponse{Providers: providers}, nil
}
// 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
@@ -751,3 +847,44 @@ func parseOptionalInt32(n *int32) sql.NullInt32 {
}
return sql.NullInt32{Int32: *n, Valid: true}
}
// aiProviderToProto maps a single ai_providers row (and its keys, for enabled
// providers) to the proto representation served to AI Gateway daemons. Keys and
// Bedrock settings are only attached for enabled providers; disabled providers
// never call upstream so their secrets are withheld.
func aiProviderToProto(row database.AIProvider, keys []database.AIProviderKey) (*proto.AIProvider, error) {
p := &proto.AIProvider{
Name: row.Name,
Type: string(row.Type),
Enabled: row.Enabled,
BaseUrl: row.BaseUrl,
}
// Disabled providers are rendered as stubs by the client and never call
// upstream, so only the identity fields are returned; keys and settings
// (including Bedrock credentials) are withheld.
if !row.Enabled {
return p, nil
}
p.Keys = make([]string, 0, len(keys))
for _, k := range keys {
p.Keys = append(p.Keys, k.APIKey)
}
settings, err := db2sdk.AIProviderSettings(row.Settings)
if err != nil {
return nil, xerrors.Errorf("decode settings: %w", err)
}
if settings.Bedrock != nil {
p.Bedrock = &proto.AIProviderKindBedrock{
Region: settings.Bedrock.Region,
AccessKey: ptr.NilToEmpty(settings.Bedrock.AccessKey),
AccessKeySecret: ptr.NilToEmpty(settings.Bedrock.AccessKeySecret),
Model: settings.Bedrock.Model,
SmallFastModel: settings.Bedrock.SmallFastModel,
RoleArn: settings.Bedrock.RoleARN,
}
}
return p, nil
}
@@ -9,6 +9,7 @@ import (
"fmt"
"net"
"net/url"
"strconv"
"testing"
"time"
@@ -25,6 +26,7 @@ import (
"cdr.dev/slog/v3"
"cdr.dev/slog/v3/sloggers/slogjson"
"cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/coderd/aibridged"
"github.com/coder/coder/v2/coderd/aibridged/proto"
"github.com/coder/coder/v2/coderd/aibridgedserver"
@@ -2405,3 +2407,207 @@ func TestInferredThreadsByToolCalls(t *testing.T) {
require.Equal(t, uuid.NullUUID{UUID: bID, Valid: true}, intcC.ThreadParentID)
require.Equal(t, uuid.NullUUID{UUID: aID, Valid: true}, intcC.ThreadRootID)
}
// TestGetAIProviders exercises the row-to-proto mapping over a real database:
// enabled providers carry their keys (and typed Bedrock settings), disabled
// providers are included but withhold keys and settings, Copilot (a keyless
// BYOK provider) round-trips with no keys, and an enabled provider whose
// settings blob cannot be decoded is skipped rather than failing the fetch.
func TestGetAIProviders(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitLong)
// The skipped misconfigured provider is logged at Error level by design,
// so error logs are expected here.
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
// Enabled OpenAI with two keys.
openai := dbgen.AIProvider(t, db, database.AIProvider{
Type: database.AIProviderTypeOpenai,
Name: "openai",
Enabled: true,
BaseUrl: "https://api.openai.com/",
})
dbgen.AIProviderKey(t, db, database.AIProviderKey{ProviderID: openai.ID, APIKey: "sk-openai-1"})
dbgen.AIProviderKey(t, db, database.AIProviderKey{ProviderID: openai.ID, APIKey: "sk-openai-2"})
// Enabled Bedrock with typed settings.
bedrockSettings, err := json.Marshal(codersdk.AIProviderSettings{
Bedrock: &codersdk.AIProviderBedrockSettings{
Region: "us-east-1",
Model: "anthropic.claude-3",
SmallFastModel: "anthropic.claude-haiku",
AccessKey: ptr.Ref("AKID"),
AccessKeySecret: ptr.Ref("secret"),
RoleARN: "arn:aws:iam::123456789012:role/bedrock",
},
})
require.NoError(t, err)
dbgen.AIProvider(t, db, database.AIProvider{
Type: database.AIProviderTypeBedrock,
Name: "bedrock",
Enabled: true,
BaseUrl: "https://bedrock-runtime.us-east-1.amazonaws.com/",
Settings: sql.NullString{String: string(bedrockSettings), Valid: true},
})
// Enabled Copilot, which is keyless (BYOK per request).
dbgen.AIProvider(t, db, database.AIProvider{
Type: database.AIProviderTypeCopilot,
Name: "copilot",
Enabled: true,
BaseUrl: "https://api.githubcopilot.com/",
})
// Disabled Anthropic with a key; the key must be withheld.
disabled := dbgen.AIProvider(t, db, database.AIProvider{
Type: database.AIProviderTypeAnthropic,
Name: "anthropic-off",
BaseUrl: "https://api.anthropic.com/",
}, func(p *database.InsertAIProviderParams) {
p.Enabled = false
})
dbgen.AIProviderKey(t, db, database.AIProviderKey{ProviderID: disabled.ID, APIKey: "sk-secret"})
// Enabled provider with an undecodable settings blob; it must be skipped
// so one corrupt row does not break provider config for every gateway.
dbgen.AIProvider(t, db, database.AIProvider{
Type: database.AIProviderTypeBedrock,
Name: "broken-settings",
Enabled: true,
BaseUrl: "https://bedrock-runtime.us-east-1.amazonaws.com/",
Settings: sql.NullString{String: "{not valid json", Valid: true},
})
srv, err := aibridgedserver.NewServer(ctx, db, logger, "/", codersdk.AIBridgeConfig{}, nil, nil, agplaiseats.Noop{})
require.NoError(t, err)
resp, err := srv.GetAIProviders(ctx, &proto.GetAIProvidersRequest{})
require.NoError(t, err)
byName := make(map[string]*proto.AIProvider, len(resp.GetProviders()))
for _, p := range resp.GetProviders() {
byName[p.GetName()] = p
}
require.Len(t, byName, 4)
assert.NotContains(t, byName, "broken-settings", "provider with undecodable settings must be skipped")
gotOpenAI := byName["openai"]
require.NotNil(t, gotOpenAI)
assert.True(t, gotOpenAI.GetEnabled())
assert.Equal(t, string(database.AIProviderTypeOpenai), gotOpenAI.GetType())
assert.Equal(t, "https://api.openai.com/", gotOpenAI.GetBaseUrl())
assert.ElementsMatch(t, []string{"sk-openai-1", "sk-openai-2"}, gotOpenAI.GetKeys())
assert.Nil(t, gotOpenAI.GetBedrock())
gotBedrock := byName["bedrock"]
require.NotNil(t, gotBedrock)
assert.True(t, gotBedrock.GetEnabled())
require.NotNil(t, gotBedrock.GetBedrock())
assert.Equal(t, "us-east-1", gotBedrock.GetBedrock().GetRegion())
assert.Equal(t, "anthropic.claude-3", gotBedrock.GetBedrock().GetModel())
assert.Equal(t, "anthropic.claude-haiku", gotBedrock.GetBedrock().GetSmallFastModel())
assert.Equal(t, "AKID", gotBedrock.GetBedrock().GetAccessKey())
assert.Equal(t, "secret", gotBedrock.GetBedrock().GetAccessKeySecret())
assert.Equal(t, "arn:aws:iam::123456789012:role/bedrock", gotBedrock.GetBedrock().GetRoleArn())
gotCopilot := byName["copilot"]
require.NotNil(t, gotCopilot)
assert.True(t, gotCopilot.GetEnabled())
assert.Empty(t, gotCopilot.GetKeys())
gotDisabled := byName["anthropic-off"]
require.NotNil(t, gotDisabled)
assert.False(t, gotDisabled.GetEnabled())
assert.Empty(t, gotDisabled.GetKeys(), "keys must be withheld for disabled providers")
assert.Nil(t, gotDisabled.GetBedrock())
}
// TestGetAIProvidersBlocksOnSeedLock asserts that GetAIProviders serializes on
// LockIDAIProvidersEnvSeed: while an in-flight seed transaction holds the lock,
// the fetch blocks, and once the seed commits the fetch returns the seeded
// set. Postgres advisory locks are required, so this cannot run against the
// mock store.
func TestGetAIProvidersBlocksOnSeedLock(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitLong)
logger := slogtest.Make(t, nil)
dbgen.AIProviderWithOptionalKey(t, db, database.AIProvider{
Type: database.AIProviderTypeOpenai,
Name: "openai",
Enabled: true,
BaseUrl: "https://api.openai.com/",
}, "sk-openai")
srv, err := aibridgedserver.NewServer(ctx, db, logger, "/", codersdk.AIBridgeConfig{}, nil, nil, agplaiseats.Noop{})
require.NoError(t, err)
// Simulate an in-flight env seed holding the advisory lock until released.
holderReady := make(chan struct{})
releaseHolder := make(chan struct{})
holderDone := make(chan struct{})
go func() {
defer close(holderDone)
txErr := db.InTx(func(tx database.Store) error {
if err := tx.AcquireLock(ctx, database.LockIDAIProvidersEnvSeed); err != nil {
return err
}
close(holderReady)
<-releaseHolder
return nil
}, nil)
assert.NoError(t, txErr)
}()
testutil.TryReceive(ctx, t, holderReady)
fetchDone := make(chan *proto.GetAIProvidersResponse, 1)
fetchErr := make(chan error, 1)
go func() {
resp, err := srv.GetAIProviders(ctx, &proto.GetAIProvidersRequest{})
fetchErr <- err
fetchDone <- resp
}()
// Wait until the fetch goroutine is observably blocked waiting on the seed
// advisory lock, rather than inferring it from a fixed delay. AcquireLock
// uses the single-bigint advisory lock form, so the waiter appears in
// pg_locks as an ungranted "advisory" row whose objid is the low 32 bits of
// the lock ID. Asserting the wait directly stops this from passing vacuously
// if the goroutine has not yet reached the lock.
require.Eventually(t, func() bool {
locks, err := db.PGLocks(ctx)
if err != nil {
return false
}
for _, l := range locks {
if l.LockType != nil && *l.LockType == "advisory" && !l.Granted &&
l.ObjID != nil && *l.ObjID == strconv.Itoa(database.LockIDAIProvidersEnvSeed) {
return true
}
}
return false
}, testutil.WaitShort, testutil.IntervalFast, "fetch must block waiting on the seed advisory lock")
// With the fetch proven to be blocked on the lock, it must not have
// completed while the lock is still held.
select {
case <-fetchDone:
t.Fatal("GetAIProviders returned before the seed lock was released")
default:
}
// Release the lock; the fetch should now complete and return the seeded set.
close(releaseHolder)
testutil.TryReceive(ctx, t, holderDone)
require.NoError(t, testutil.TryReceive(ctx, t, fetchErr))
resp := testutil.TryReceive(ctx, t, fetchDone)
require.Len(t, resp.GetProviders(), 1)
assert.Equal(t, "openai", resp.GetProviders()[0].GetName())
assert.Equal(t, []string{"sk-openai"}, resp.GetProviders()[0].GetKeys())
}
+7 -4
View File
@@ -7,10 +7,10 @@ import (
"github.com/coder/coder/v2/coderd/aibridged/proto"
)
// Register registers the Recorder, MCPConfigurator, and Authorizer DRPC
// services backed by srv onto mux. It is shared by the embedded in-memory
// server and the standalone /api/v2/ai-gateway/serve WebSocket handler so both
// expose an identical service set.
// Register registers the Recorder, MCPConfigurator, Authorizer, and
// ProviderConfigurator DRPC services backed by srv onto mux. It is shared by
// the embedded in-memory server and the standalone /api/v2/ai-gateway/serve
// WebSocket handler so both expose an identical service set.
func Register(mux *drpcmux.Mux, srv *Server) error {
if err := proto.DRPCRegisterRecorder(mux, srv); err != nil {
return xerrors.Errorf("register recorder service: %w", err)
@@ -21,5 +21,8 @@ func Register(mux *drpcmux.Mux, srv *Server) error {
if err := proto.DRPCRegisterAuthorizer(mux, srv); err != nil {
return xerrors.Errorf("register authorizer service: %w", err)
}
if err := proto.DRPCRegisterProviderConfigurator(mux, srv); err != nil {
return xerrors.Errorf("register provider configurator service: %w", err)
}
return nil
}