mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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,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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user