feat: use AI provider chat APIs (#25415)

This commit is contained in:
Michael Suchacz
2026-05-22 07:53:23 +02:00
committed by GitHub
parent 10efde3e6c
commit 06526a5822
41 changed files with 2195 additions and 1126 deletions
+40
View File
@@ -157,6 +157,46 @@ func TestAIProvidersCRUD(t *testing.T) {
require.Equal(t, "no-display", created.DisplayName)
})
t.Run("RequiredBaseURL", func(t *testing.T) {
t.Parallel()
client := coderdtest.New(t, nil)
_ = coderdtest.CreateFirstUser(t, client)
ctx := testutil.Context(t, testutil.WaitLong)
//nolint:gocritic // Owner role is the audience for this endpoint.
_, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderTypeOpenAI,
Name: "missing-base-url",
Enabled: true,
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "Invalid AI provider request.", sdkErr.Message)
require.Contains(t, sdkErr.Validations, codersdk.ValidationError{Field: "base_url", Detail: "base_url is required"})
created, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderTypeOpenAI,
Name: "required-base-url",
Enabled: true,
BaseURL: "https://api.openai.com/v1",
})
require.NoError(t, err)
baseURL := "https://proxy.example.com/v1"
updated, err := client.UpdateAIProvider(ctx, created.Name, codersdk.UpdateAIProviderRequest{
BaseURL: &baseURL,
})
require.NoError(t, err)
require.Equal(t, baseURL, updated.BaseURL)
baseURL = ""
_, err = client.UpdateAIProvider(ctx, created.Name, codersdk.UpdateAIProviderRequest{
BaseURL: &baseURL,
})
sdkErr = requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "Invalid AI provider request.", sdkErr.Message)
require.Contains(t, sdkErr.Validations, codersdk.ValidationError{Field: "base_url", Detail: "base_url is required"})
})
t.Run("DuplicateNameConflict", func(t *testing.T) {
t.Parallel()
client := coderdtest.New(t, nil)
+11
View File
@@ -1202,6 +1202,17 @@ func New(options *Options) *API {
r.Delete("/", api.deleteUserSkill)
})
})
r.Route("/users/{user}/ai-provider-keys", func(r chi.Router) {
r.Use(
apiKeyMiddleware,
httpmw.ExtractUserParam(options.Database),
)
r.Get("/", api.listUserAIProviderKeyConfigs)
r.Route("/{aiProvider}", func(r chi.Router) {
r.Put("/", api.upsertUserAIProviderKey)
r.Delete("/", api.deleteUserAIProviderKey)
})
})
r.Route("/chats", func(r chi.Router) {
r.Use(
apiKeyMiddleware,
+13
View File
@@ -65,11 +65,24 @@ func CreateOpenAICompatChatModelConfig(
BaseURL: baseURL,
})
require.NoError(t, err)
aiProviderBaseURL := baseURL
if aiProviderBaseURL == "" {
aiProviderBaseURL = "https://api.example.com/v1"
}
provider, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderType(TestChatProviderOpenAICompat),
Name: "test-" + uuid.NewString(),
BaseURL: aiProviderBaseURL,
Enabled: true,
APIKeys: []string{TestChatProviderAPIKey},
})
require.NoError(t, err)
contextLimit := int64(4096)
isDefault := true
modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: TestChatProviderOpenAICompat,
AIProviderID: &provider.ID,
Model: TestChatModelOpenAICompat,
ContextLimit: &contextLimit,
IsDefault: &isDefault,
+38
View File
@@ -722,6 +722,24 @@ var (
}),
Scope: rbac.ScopeAll,
}.WithCachedASTValue()
subjectAIProviderMetadataReader = rbac.Subject{
Type: rbac.SubjectTypeAIProviderMetadataReader,
FriendlyName: "AI Provider Metadata Reader",
ID: uuid.Nil.String(),
Roles: rbac.Roles([]rbac.Role{
{
Identifier: rbac.RoleIdentifier{Name: "ai-provider-metadata-reader"},
DisplayName: "AI Provider Metadata Reader",
Site: rbac.Permissions(map[string][]policy.Action{
rbac.ResourceAIProvider.Type: {policy.ActionRead},
}),
User: []rbac.Permission{},
ByOrgID: map[string]rbac.OrgPermissions{},
},
}),
Scope: rbac.ScopeAll,
}.WithCachedASTValue()
)
// AsProvisionerd returns a context with an actor that has permissions required
@@ -846,6 +864,12 @@ func AsChatd(ctx context.Context) context.Context {
return As(ctx, subjectChatd)
}
// AsAIProviderMetadataReader returns a context with an actor that can read
// AI provider metadata and provider-key presence.
func AsAIProviderMetadataReader(ctx context.Context) context.Context {
return As(ctx, subjectAIProviderMetadataReader)
}
var AsRemoveActor = rbac.Subject{
ID: "remove-actor",
}
@@ -2546,6 +2570,13 @@ func (q *querier) GetAIProviderByID(ctx context.Context, id uuid.UUID) (database
return q.db.GetAIProviderByID(ctx, id)
}
func (q *querier) GetAIProviderByIDForReferenceLock(ctx context.Context, id uuid.UUID) (database.AIProvider, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceAIProvider); err != nil {
return database.AIProvider{}, err
}
return q.db.GetAIProviderByIDForReferenceLock(ctx, id)
}
func (q *querier) GetAIProviderByName(ctx context.Context, name string) (database.AIProvider, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceAIProvider); err != nil {
return database.AIProvider{}, err
@@ -2560,6 +2591,13 @@ func (q *querier) GetAIProviderKeyByID(ctx context.Context, id uuid.UUID) (datab
return q.db.GetAIProviderKeyByID(ctx, id)
}
func (q *querier) GetAIProviderKeyPresence(ctx context.Context, arg []uuid.UUID) ([]uuid.UUID, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceAIProvider); err != nil {
return nil, err
}
return q.db.GetAIProviderKeyPresence(ctx, arg)
}
func (q *querier) GetAIProviderKeys(ctx context.Context, includeDeleted bool) ([]database.AIProviderKey, error) {
// Callers pass include_deleted=TRUE only from the dbcrypt key
// rotation utility, which needs to re-encrypt every row that holds
+13
View File
@@ -6509,6 +6509,11 @@ func (s *MethodTestSuite) TestAIBridge() {
dbm.EXPECT().GetAIProviderByID(gomock.Any(), provider.ID).Return(provider, nil).AnyTimes()
check.Args(provider.ID).Asserts(rbac.ResourceAIProvider, policy.ActionRead).Returns(provider)
}))
s.Run("GetAIProviderByIDForReferenceLock", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
provider := testutil.Fake(s.T(), faker, database.AIProvider{})
dbm.EXPECT().GetAIProviderByIDForReferenceLock(gomock.Any(), provider.ID).Return(provider, nil).AnyTimes()
check.Args(provider.ID).Asserts(rbac.ResourceAIProvider, policy.ActionRead).Returns(provider)
}))
s.Run("GetAIProviderByName", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
provider := testutil.Fake(s.T(), faker, database.AIProvider{})
dbm.EXPECT().GetAIProviderByName(gomock.Any(), provider.Name).Return(provider, nil).AnyTimes()
@@ -6562,6 +6567,14 @@ func (s *MethodTestSuite) TestAIBridge() {
dbm.EXPECT().GetAIProviderKeyByID(gomock.Any(), key.ID).Return(key, nil).AnyTimes()
check.Args(key.ID).Asserts(rbac.ResourceAIProvider, policy.ActionRead).Returns(key)
}))
s.Run("GetAIProviderKeyPresence", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
providerA := testutil.Fake(s.T(), faker, database.AIProvider{})
providerB := testutil.Fake(s.T(), faker, database.AIProvider{})
arg := []uuid.UUID{providerA.ID, providerB.ID}
providerIDs := []uuid.UUID{providerA.ID}
dbm.EXPECT().GetAIProviderKeyPresence(gomock.Any(), arg).Return(providerIDs, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceAIProvider, policy.ActionRead).Returns(providerIDs)
}))
s.Run("GetAIProviderKeysByProviderID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
provider := testutil.Fake(s.T(), faker, database.AIProvider{})
keyA := testutil.Fake(s.T(), faker, database.AIProviderKey{ProviderID: provider.ID})
+1
View File
@@ -160,6 +160,7 @@ func ChatModelConfig(t testing.TB, db database.Store, seed database.ChatModelCon
ContextLimit: takeFirst(seed.ContextLimit, defaultChatModelContextLimit),
CompressionThreshold: takeFirst(seed.CompressionThreshold, defaultChatModelCompressionThreshold),
Options: takeFirstSlice(seed.Options, json.RawMessage(`{}`)),
AIProviderID: seed.AIProviderID,
}
for _, fn := range munge {
fn(&params)
+16
View File
@@ -1041,6 +1041,14 @@ func (m queryMetricsStore) GetAIProviderByID(ctx context.Context, id uuid.UUID)
return r0, r1
}
func (m queryMetricsStore) GetAIProviderByIDForReferenceLock(ctx context.Context, id uuid.UUID) (database.AIProvider, error) {
start := time.Now()
r0, r1 := m.s.GetAIProviderByIDForReferenceLock(ctx, id)
m.queryLatencies.WithLabelValues("GetAIProviderByIDForReferenceLock").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetAIProviderByIDForReferenceLock").Inc()
return r0, r1
}
func (m queryMetricsStore) GetAIProviderByName(ctx context.Context, name string) (database.AIProvider, error) {
start := time.Now()
r0, r1 := m.s.GetAIProviderByName(ctx, name)
@@ -1057,6 +1065,14 @@ func (m queryMetricsStore) GetAIProviderKeyByID(ctx context.Context, id uuid.UUI
return r0, r1
}
func (m queryMetricsStore) GetAIProviderKeyPresence(ctx context.Context, arg []uuid.UUID) ([]uuid.UUID, error) {
start := time.Now()
r0, r1 := m.s.GetAIProviderKeyPresence(ctx, arg)
m.queryLatencies.WithLabelValues("GetAIProviderKeyPresence").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetAIProviderKeyPresence").Inc()
return r0, r1
}
func (m queryMetricsStore) GetAIProviderKeys(ctx context.Context, includeDeleted bool) ([]database.AIProviderKey, error) {
start := time.Now()
r0, r1 := m.s.GetAIProviderKeys(ctx, includeDeleted)
+30
View File
@@ -1799,6 +1799,21 @@ func (mr *MockStoreMockRecorder) GetAIProviderByID(ctx, id any) *gomock.Call {
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAIProviderByID", reflect.TypeOf((*MockStore)(nil).GetAIProviderByID), ctx, id)
}
// GetAIProviderByIDForReferenceLock mocks base method.
func (m *MockStore) GetAIProviderByIDForReferenceLock(ctx context.Context, id uuid.UUID) (database.AIProvider, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetAIProviderByIDForReferenceLock", ctx, id)
ret0, _ := ret[0].(database.AIProvider)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetAIProviderByIDForReferenceLock indicates an expected call of GetAIProviderByIDForReferenceLock.
func (mr *MockStoreMockRecorder) GetAIProviderByIDForReferenceLock(ctx, id any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAIProviderByIDForReferenceLock", reflect.TypeOf((*MockStore)(nil).GetAIProviderByIDForReferenceLock), ctx, id)
}
// GetAIProviderByName mocks base method.
func (m *MockStore) GetAIProviderByName(ctx context.Context, name string) (database.AIProvider, error) {
m.ctrl.T.Helper()
@@ -1829,6 +1844,21 @@ func (mr *MockStoreMockRecorder) GetAIProviderKeyByID(ctx, id any) *gomock.Call
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAIProviderKeyByID", reflect.TypeOf((*MockStore)(nil).GetAIProviderKeyByID), ctx, id)
}
// GetAIProviderKeyPresence mocks base method.
func (m *MockStore) GetAIProviderKeyPresence(ctx context.Context, providerIds []uuid.UUID) ([]uuid.UUID, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetAIProviderKeyPresence", ctx, providerIds)
ret0, _ := ret[0].([]uuid.UUID)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetAIProviderKeyPresence indicates an expected call of GetAIProviderKeyPresence.
func (mr *MockStoreMockRecorder) GetAIProviderKeyPresence(ctx, providerIds any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAIProviderKeyPresence", reflect.TypeOf((*MockStore)(nil).GetAIProviderKeyPresence), ctx, providerIds)
}
// GetAIProviderKeys mocks base method.
func (m *MockStore) GetAIProviderKeys(ctx context.Context, includeDeleted bool) ([]database.AIProviderKey, error) {
m.ctrl.T.Helper()
+6
View File
@@ -253,8 +253,14 @@ type sqlcQuerier interface {
GetAIBridgeUserPromptsByInterceptionID(ctx context.Context, interceptionID uuid.UUID) ([]AIBridgeUserPrompt, error)
GetAIModelPriceByProviderModel(ctx context.Context, arg GetAIModelPriceByProviderModelParams) (AiModelPrice, error)
GetAIProviderByID(ctx context.Context, id uuid.UUID) (AIProvider, error)
// Lock the provider row until the model-config write completes. The
// transaction alone does not stop a concurrent soft-delete or disable
// between validation and writing the model config reference.
GetAIProviderByIDForReferenceLock(ctx context.Context, id uuid.UUID) (AIProvider, error)
GetAIProviderByName(ctx context.Context, name string) (AIProvider, error)
GetAIProviderKeyByID(ctx context.Context, id uuid.UUID) (AIProviderKey, error)
// Returns the provider IDs that have at least one provider-scoped key.
GetAIProviderKeyPresence(ctx context.Context, providerIds []uuid.UUID) ([]uuid.UUID, error)
// Returns AI provider key rows. By default, only rows whose parent
// provider is live (deleted = FALSE) are returned, so the API list
// handler can fetch every visible provider's keys in a single query.
+71
View File
@@ -10567,6 +10567,77 @@ func TestInsertWorkspaceAgentDevcontainers(t *testing.T) {
}
}
func TestGetEnabledChatModelConfigsUsesAIProviders(t *testing.T) {
t.Parallel()
store, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitMedium)
enabledProvider := dbgen.AIProvider(t, store, database.AIProvider{
Type: database.AiProviderTypeOpenrouter,
Name: "openrouter-" + uuid.NewString(),
})
disabledProvider := dbgen.AIProvider(t, store, database.AIProvider{
Type: database.AiProviderTypeVercel,
Name: "vercel-" + uuid.NewString(),
}, func(params *database.InsertAIProviderParams) {
params.Enabled = false
})
enabledConfig := dbgen.ChatModelConfig(t, store, database.ChatModelConfig{
Provider: string(enabledProvider.Type),
Model: "openrouter-model-" + uuid.NewString(),
AIProviderID: uuid.NullUUID{
UUID: enabledProvider.ID,
Valid: true,
},
})
disabledProviderConfig := dbgen.ChatModelConfig(t, store, database.ChatModelConfig{
Provider: string(disabledProvider.Type),
Model: "vercel-model-" + uuid.NewString(),
AIProviderID: uuid.NullUUID{
UUID: disabledProvider.ID,
Valid: true,
},
})
disabledModelConfig := dbgen.ChatModelConfig(t, store, database.ChatModelConfig{
Provider: string(enabledProvider.Type),
Model: "disabled-model-" + uuid.NewString(),
AIProviderID: uuid.NullUUID{
UUID: enabledProvider.ID,
Valid: true,
},
}, func(params *database.InsertChatModelConfigParams) {
params.Enabled = false
})
legacyProvider := dbgen.ChatProvider(t, store, database.ChatProvider{Provider: "google"})
legacyConfig := dbgen.ChatModelConfig(t, store, database.ChatModelConfig{
Provider: legacyProvider.Provider,
Model: "google-model-" + uuid.NewString(),
})
configs, err := store.GetEnabledChatModelConfigs(ctx)
require.NoError(t, err)
require.True(t, slices.ContainsFunc(configs, func(config database.ChatModelConfig) bool {
return config.ID == enabledConfig.ID
}))
require.True(t, slices.ContainsFunc(configs, func(config database.ChatModelConfig) bool {
return config.ID == legacyConfig.ID
}))
require.False(t, slices.ContainsFunc(configs, func(config database.ChatModelConfig) bool {
return config.ID == disabledProviderConfig.ID
}))
require.False(t, slices.ContainsFunc(configs, func(config database.ChatModelConfig) bool {
return config.ID == disabledModelConfig.ID
}))
config, err := store.GetEnabledChatModelConfigByID(ctx, enabledConfig.ID)
require.NoError(t, err)
require.Equal(t, enabledConfig.ID, config.ID)
_, err = store.GetEnabledChatModelConfigByID(ctx, disabledProviderConfig.ID)
require.ErrorIs(t, err, sql.ErrNoRows)
}
func TestInsertChatMessages(t *testing.T) {
t.Parallel()
+93 -9
View File
@@ -146,6 +146,41 @@ func (q *sqlQuerier) GetAIProviderKeyByID(ctx context.Context, id uuid.UUID) (AI
return i, err
}
const getAIProviderKeyPresence = `-- name: GetAIProviderKeyPresence :many
SELECT DISTINCT
provider_id
FROM
ai_provider_keys
WHERE
provider_id = ANY($1::uuid[])
ORDER BY
provider_id ASC
`
// Returns the provider IDs that have at least one provider-scoped key.
func (q *sqlQuerier) GetAIProviderKeyPresence(ctx context.Context, providerIds []uuid.UUID) ([]uuid.UUID, error) {
rows, err := q.db.QueryContext(ctx, getAIProviderKeyPresence, pq.Array(providerIds))
if err != nil {
return nil, err
}
defer rows.Close()
var items []uuid.UUID
for rows.Next() {
var provider_id uuid.UUID
if err := rows.Scan(&provider_id); err != nil {
return nil, err
}
items = append(items, provider_id)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const getAIProviderKeys = `-- name: GetAIProviderKeys :many
SELECT
ai_provider_keys.id, ai_provider_keys.provider_id, ai_provider_keys.api_key, ai_provider_keys.api_key_key_id, ai_provider_keys.created_at, ai_provider_keys.updated_at
@@ -371,6 +406,38 @@ func (q *sqlQuerier) GetAIProviderByID(ctx context.Context, id uuid.UUID) (AIPro
return i, err
}
const getAIProviderByIDForReferenceLock = `-- name: GetAIProviderByIDForReferenceLock :one
SELECT
id, type, name, display_name, enabled, deleted, base_url, settings, settings_key_id, created_at, updated_at
FROM
ai_providers
WHERE
id = $1::uuid AND deleted = FALSE
FOR SHARE
`
// Lock the provider row until the model-config write completes. The
// transaction alone does not stop a concurrent soft-delete or disable
// between validation and writing the model config reference.
func (q *sqlQuerier) GetAIProviderByIDForReferenceLock(ctx context.Context, id uuid.UUID) (AIProvider, error) {
row := q.db.QueryRowContext(ctx, getAIProviderByIDForReferenceLock, id)
var i AIProvider
err := row.Scan(
&i.ID,
&i.Type,
&i.Name,
&i.DisplayName,
&i.Enabled,
&i.Deleted,
&i.BaseUrl,
&i.Settings,
&i.SettingsKeyID,
&i.CreatedAt,
&i.UpdatedAt,
)
return i, err
}
const getAIProviderByName = `-- name: GetAIProviderByName :one
SELECT
id, type, name, display_name, enabled, deleted, base_url, settings, settings_key_id, created_at, updated_at
@@ -5097,13 +5164,18 @@ SELECT
cmc.id, cmc.provider, cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, cmc.ai_provider_id
FROM
chat_model_configs cmc
JOIN
chat_providers cp ON cp.provider = cmc.provider
LEFT JOIN
ai_providers ap ON ap.id = cmc.ai_provider_id
LEFT JOIN
chat_providers cp ON cp.provider = cmc.provider AND cmc.ai_provider_id IS NULL
WHERE
cmc.id = $1::uuid
AND cmc.deleted = FALSE
AND cmc.enabled = TRUE
AND cp.enabled = TRUE
AND (
(cmc.ai_provider_id IS NOT NULL AND ap.enabled = TRUE AND ap.deleted = FALSE)
OR (cmc.ai_provider_id IS NULL AND cp.enabled = TRUE)
)
`
// Providers can be disabled independently of their model configs.
@@ -5137,12 +5209,17 @@ SELECT
cmc.id, cmc.provider, cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, cmc.ai_provider_id
FROM
chat_model_configs cmc
JOIN
chat_providers cp ON cp.provider = cmc.provider
LEFT JOIN
ai_providers ap ON ap.id = cmc.ai_provider_id
LEFT JOIN
chat_providers cp ON cp.provider = cmc.provider AND cmc.ai_provider_id IS NULL
WHERE
cmc.enabled = TRUE
AND cmc.deleted = FALSE
AND cp.enabled = TRUE
AND (
(cmc.ai_provider_id IS NOT NULL AND ap.enabled = TRUE AND ap.deleted = FALSE)
OR (cmc.ai_provider_id IS NULL AND cp.enabled = TRUE)
)
ORDER BY
cmc.provider ASC,
cmc.model ASC,
@@ -5201,7 +5278,8 @@ INSERT INTO chat_model_configs (
is_default,
context_limit,
compression_threshold,
options
options,
ai_provider_id
) VALUES (
$1::text,
$2::text,
@@ -5212,7 +5290,8 @@ INSERT INTO chat_model_configs (
$7::boolean,
$8::bigint,
$9::integer,
$10::jsonb
$10::jsonb,
$11::uuid
)
RETURNING
id, provider, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id
@@ -5229,6 +5308,7 @@ type InsertChatModelConfigParams struct {
ContextLimit int64 `db:"context_limit" json:"context_limit"`
CompressionThreshold int32 `db:"compression_threshold" json:"compression_threshold"`
Options json.RawMessage `db:"options" json:"options"`
AIProviderID uuid.NullUUID `db:"ai_provider_id" json:"ai_provider_id"`
}
func (q *sqlQuerier) InsertChatModelConfig(ctx context.Context, arg InsertChatModelConfigParams) (ChatModelConfig, error) {
@@ -5243,6 +5323,7 @@ func (q *sqlQuerier) InsertChatModelConfig(ctx context.Context, arg InsertChatMo
arg.ContextLimit,
arg.CompressionThreshold,
arg.Options,
arg.AIProviderID,
)
var i ChatModelConfig
err := row.Scan(
@@ -5295,9 +5376,10 @@ SET
context_limit = $7::bigint,
compression_threshold = $8::integer,
options = $9::jsonb,
ai_provider_id = $10::uuid,
updated_at = NOW()
WHERE
id = $10::uuid
id = $11::uuid
AND deleted = FALSE
RETURNING
id, provider, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id
@@ -5313,6 +5395,7 @@ type UpdateChatModelConfigParams struct {
ContextLimit int64 `db:"context_limit" json:"context_limit"`
CompressionThreshold int32 `db:"compression_threshold" json:"compression_threshold"`
Options json.RawMessage `db:"options" json:"options"`
AIProviderID uuid.NullUUID `db:"ai_provider_id" json:"ai_provider_id"`
ID uuid.UUID `db:"id" json:"id"`
}
@@ -5327,6 +5410,7 @@ func (q *sqlQuerier) UpdateChatModelConfig(ctx context.Context, arg UpdateChatMo
arg.ContextLimit,
arg.CompressionThreshold,
arg.Options,
arg.AIProviderID,
arg.ID,
)
var i ChatModelConfig
@@ -21,6 +21,17 @@ ORDER BY
created_at ASC,
id ASC;
-- name: GetAIProviderKeyPresence :many
-- Returns the provider IDs that have at least one provider-scoped key.
SELECT DISTINCT
provider_id
FROM
ai_provider_keys
WHERE
provider_id = ANY(@provider_ids::uuid[])
ORDER BY
provider_id ASC;
-- name: GetAIProviderKeys :many
-- Returns AI provider key rows. By default, only rows whose parent
-- provider is live (deleted = FALSE) are returned, so the API list
+12
View File
@@ -6,6 +6,18 @@ FROM
WHERE
id = @id::uuid AND deleted = FALSE;
-- name: GetAIProviderByIDForReferenceLock :one
SELECT
*
FROM
ai_providers
WHERE
id = @id::uuid AND deleted = FALSE
-- Lock the provider row until the model-config write completes. The
-- transaction alone does not stop a concurrent soft-delete or disable
-- between validation and writing the model config reference.
FOR SHARE;
-- name: GetAIProviderByName :one
SELECT
*
+21 -8
View File
@@ -34,12 +34,17 @@ SELECT
cmc.*
FROM
chat_model_configs cmc
JOIN
chat_providers cp ON cp.provider = cmc.provider
LEFT JOIN
ai_providers ap ON ap.id = cmc.ai_provider_id
LEFT JOIN
chat_providers cp ON cp.provider = cmc.provider AND cmc.ai_provider_id IS NULL
WHERE
cmc.enabled = TRUE
AND cmc.deleted = FALSE
AND cp.enabled = TRUE
AND (
(cmc.ai_provider_id IS NOT NULL AND ap.enabled = TRUE AND ap.deleted = FALSE)
OR (cmc.ai_provider_id IS NULL AND cp.enabled = TRUE)
)
ORDER BY
cmc.provider ASC,
cmc.model ASC,
@@ -53,13 +58,18 @@ FROM
chat_model_configs cmc
-- Providers can be disabled independently of their model configs.
-- Check both to ensure the selected config is actually usable.
JOIN
chat_providers cp ON cp.provider = cmc.provider
LEFT JOIN
ai_providers ap ON ap.id = cmc.ai_provider_id
LEFT JOIN
chat_providers cp ON cp.provider = cmc.provider AND cmc.ai_provider_id IS NULL
WHERE
cmc.id = @id::uuid
AND cmc.deleted = FALSE
AND cmc.enabled = TRUE
AND cp.enabled = TRUE;
AND (
(cmc.ai_provider_id IS NOT NULL AND ap.enabled = TRUE AND ap.deleted = FALSE)
OR (cmc.ai_provider_id IS NULL AND cp.enabled = TRUE)
);
-- name: InsertChatModelConfig :one
INSERT INTO chat_model_configs (
@@ -72,7 +82,8 @@ INSERT INTO chat_model_configs (
is_default,
context_limit,
compression_threshold,
options
options,
ai_provider_id
) VALUES (
@provider::text,
@model::text,
@@ -83,7 +94,8 @@ INSERT INTO chat_model_configs (
@is_default::boolean,
@context_limit::bigint,
@compression_threshold::integer,
@options::jsonb
@options::jsonb,
sqlc.narg('ai_provider_id')::uuid
)
RETURNING
*;
@@ -101,6 +113,7 @@ SET
context_limit = @context_limit::bigint,
compression_threshold = @compression_threshold::integer,
options = @options::jsonb,
ai_provider_id = sqlc.narg('ai_provider_id')::uuid,
updated_at = NOW()
WHERE
id = @id::uuid
+265 -16
View File
@@ -6396,6 +6396,179 @@ func convertChatMessages(messages []database.ChatMessage) []codersdk.ChatMessage
return result
}
func parseUserAIProviderID(r *http.Request) (uuid.UUID, error) {
return uuid.Parse(chi.URLParam(r, "aiProvider"))
}
func convertAIProviderSummary(provider database.AIProvider) codersdk.AIProviderSummary {
displayName := provider.Name
if provider.DisplayName.Valid && provider.DisplayName.String != "" {
displayName = provider.DisplayName.String
}
return codersdk.AIProviderSummary{
ID: provider.ID,
Type: codersdk.AIProviderType(provider.Type),
Name: provider.Name,
DisplayName: displayName,
Enabled: provider.Enabled,
Deleted: provider.Deleted,
}
}
func (api *API) listUserAIProviderKeyConfigs(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
targetUser := httpmw.UserParam(r)
//nolint:gocritic // Users can list limited provider metadata to manage their own AI provider keys.
metadataCtx := dbauthz.AsAIProviderMetadataReader(ctx)
providers, err := api.Database.GetAIProviders(metadataCtx, database.GetAIProvidersParams{IncludeDisabled: true})
if err != nil {
api.Logger.Error(ctx, "failed to list user AI provider configs", slog.Error(err), slog.F("user_id", targetUser.ID))
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{Message: "Failed to list AI providers."})
return
}
keys, err := api.Database.GetUserAIProviderKeysByUserID(ctx, targetUser.ID)
if err != nil {
api.Logger.Error(ctx, "failed to list user AI provider keys", slog.Error(err), slog.F("user_id", targetUser.ID))
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{Message: "Failed to list user AI provider keys."})
return
}
keysByProviderID := make(map[uuid.UUID]struct{}, len(keys))
for _, key := range keys {
keysByProviderID[key.AIProviderID] = struct{}{}
}
visibleProviders := make([]database.AIProvider, 0, len(providers))
visibleProviderIDs := make([]uuid.UUID, 0, len(providers))
for _, provider := range providers {
_, hasUserKey := keysByProviderID[provider.ID]
if !provider.Enabled && !hasUserKey {
continue
}
visibleProviders = append(visibleProviders, provider)
visibleProviderIDs = append(visibleProviderIDs, provider.ID)
}
providerKeysByProviderID := make(map[uuid.UUID]struct{}, len(visibleProviderIDs))
if len(visibleProviderIDs) > 0 {
providerKeyIDs, err := api.Database.GetAIProviderKeyPresence(metadataCtx, visibleProviderIDs)
if err != nil {
api.Logger.Error(ctx, "failed to list AI provider key presence", slog.Error(err), slog.F("user_id", targetUser.ID))
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{Message: "Failed to list AI provider keys."})
return
}
for _, providerID := range providerKeyIDs {
providerKeysByProviderID[providerID] = struct{}{}
}
}
byokEnabled := api.DeploymentValues.AI.BridgeConfig.AllowBYOK.Value()
configs := make([]codersdk.UserAIProviderKeyConfig, 0, len(visibleProviders))
for _, provider := range visibleProviders {
_, hasUserKey := keysByProviderID[provider.ID]
_, hasProviderKey := providerKeysByProviderID[provider.ID]
configs = append(configs, codersdk.UserAIProviderKeyConfig{
Provider: convertAIProviderSummary(provider),
HasUserAPIKey: hasUserKey,
HasProviderAPIKey: hasProviderKey,
BYOKEnabled: byokEnabled,
})
}
httpapi.Write(ctx, rw, http.StatusOK, configs)
}
func (api *API) upsertUserAIProviderKey(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
if !api.DeploymentValues.AI.BridgeConfig.AllowBYOK.Value() {
httpapi.Write(ctx, rw, http.StatusForbidden, codersdk.Response{Message: "BYOK is disabled."})
return
}
targetUser := httpmw.UserParam(r)
providerID, err := parseUserAIProviderID(r)
if err != nil {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{Message: "Invalid AI provider ID."})
return
}
//nolint:gocritic // Users can attach their own key to an enabled provider without AI provider admin permissions.
metadataCtx := dbauthz.AsAIProviderMetadataReader(ctx)
provider, err := api.Database.GetAIProviderByID(metadataCtx, providerID)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
httpapi.Write(ctx, rw, http.StatusNotFound, codersdk.Response{Message: "AI provider not found."})
return
}
api.Logger.Error(ctx, "failed to get AI provider", slog.Error(err), slog.F("ai_provider_id", providerID))
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{Message: "Failed to get AI provider."})
return
}
if !provider.Enabled {
httpapi.Write(ctx, rw, http.StatusPreconditionFailed, codersdk.Response{Message: "AI provider is disabled."})
return
}
var req codersdk.CreateUserAIProviderKeyRequest
if !httpapi.Read(ctx, rw, r, &req) {
return
}
if err := validateChatProviderAPIKeySize(req.APIKey); err != nil {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "API key too large.",
Detail: err.Error(),
})
return
}
if req.APIKey == "" {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{Message: "API key is required."})
return
}
if strings.TrimSpace(req.APIKey) != req.APIKey {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{Message: "API key must not contain leading or trailing whitespace."})
return
}
providerKeys, err := api.Database.GetAIProviderKeyPresence(metadataCtx, []uuid.UUID{providerID})
if err != nil {
api.Logger.Error(ctx, "failed to list AI provider key presence", slog.Error(err), slog.F("ai_provider_id", providerID))
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{Message: "Failed to list AI provider keys."})
return
}
now := api.Clock.Now()
_, err = api.Database.UpsertUserAIProviderKey(ctx, database.UpsertUserAIProviderKeyParams{
ID: uuid.New(),
UserID: targetUser.ID,
AIProviderID: providerID,
APIKey: req.APIKey,
ApiKeyKeyID: sql.NullString{},
CreatedAt: now,
UpdatedAt: now,
})
if err != nil {
api.Logger.Error(ctx, "failed to update user AI provider key", slog.Error(err), slog.F("user_id", targetUser.ID), slog.F("ai_provider_id", providerID))
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{Message: "Failed to update user AI provider key."})
return
}
httpapi.Write(ctx, rw, http.StatusOK, codersdk.UserAIProviderKeyConfig{
Provider: convertAIProviderSummary(provider),
HasUserAPIKey: true,
HasProviderAPIKey: len(providerKeys) > 0,
BYOKEnabled: true,
})
}
func (api *API) deleteUserAIProviderKey(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
targetUser := httpmw.UserParam(r)
providerID, err := parseUserAIProviderID(r)
if err != nil {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{Message: "Invalid AI provider ID."})
return
}
if err := api.Database.DeleteUserAIProviderKey(ctx, database.DeleteUserAIProviderKeyParams{UserID: targetUser.ID, AIProviderID: providerID}); err != nil {
api.Logger.Error(ctx, "failed to delete user AI provider key", slog.Error(err), slog.F("user_id", targetUser.ID), slog.F("ai_provider_id", providerID))
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{Message: "Failed to delete user AI provider key."})
return
}
httpapi.Write(ctx, rw, http.StatusNoContent, nil)
}
func (api *API) listChatProviders(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
//nolint:gocritic // System context required to read enabled chat providers.
@@ -6890,6 +7063,7 @@ func (api *API) listUserChatProviderConfigs(rw http.ResponseWriter, r *http.Requ
provider,
hasUserAPIKey,
hasCentralAPIKeyFallback,
api.DeploymentValues.AI.BridgeConfig.AllowBYOK.Value(),
),
)
}
@@ -6978,6 +7152,7 @@ func (api *API) upsertUserChatProviderKey(rw http.ResponseWriter, r *http.Reques
provider,
true,
hasCentralAPIKeyFallback,
api.DeploymentValues.AI.BridgeConfig.AllowBYOK.Value(),
),
)
}
@@ -7052,14 +7227,29 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) {
return
}
provider := normalizeChatProvider(req.Provider)
if provider == "" {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "Invalid provider.",
Detail: chatProviderValidationDetail(),
if req.AIProviderID == nil {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{Message: "AI provider ID is required."})
return
}
//nolint:gocritic // The route already authorized chat model config updates.
aiProvider, err := api.Database.GetAIProviderByID(dbauthz.AsChatd(ctx), *req.AIProviderID)
if err != nil {
if httpapi.Is404Error(err) {
httpapi.Write(ctx, rw, http.StatusPreconditionFailed, codersdk.Response{Message: "AI provider is not configured."})
return
}
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to get AI provider.",
Detail: err.Error(),
})
return
}
if !aiProvider.Enabled {
httpapi.Write(ctx, rw, http.StatusPreconditionFailed, codersdk.Response{Message: "AI provider is disabled."})
return
}
provider := string(aiProvider.Type)
aiProviderID := uuid.NullUUID{UUID: aiProvider.ID, Valid: true}
model := strings.TrimSpace(req.Model)
if model == "" {
@@ -7117,15 +7307,25 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) {
ContextLimit: contextLimit,
CompressionThreshold: compressionThreshold,
Options: modelConfigRaw,
AIProviderID: aiProviderID,
CreatedBy: uuid.NullUUID{UUID: apiKey.UserID, Valid: apiKey.UserID != uuid.Nil},
UpdatedBy: uuid.NullUUID{UUID: apiKey.UserID, Valid: apiKey.UserID != uuid.Nil},
}
var inserted database.ChatModelConfig
err := api.Database.InTx(func(tx database.Store) error {
if err := requireChatProviderForModelConfig(ctx, tx, insertParams.Provider); err != nil {
return err
err = api.Database.InTx(func(tx database.Store) error {
//nolint:gocritic // The route already authorized chat model config updates.
lockedAIProvider, err := tx.GetAIProviderByIDForReferenceLock(dbauthz.AsChatd(ctx), insertParams.AIProviderID.UUID)
if err != nil {
if xerrors.Is(err, sql.ErrNoRows) {
return errChatProviderNotConfigured
}
return xerrors.Errorf("get AI provider for update: %w", err)
}
if !lockedAIProvider.Enabled {
return errChatProviderNotConfigured
}
insertParams.Provider = string(lockedAIProvider.Type)
insertAsDefault := isDefault
if !insertAsDefault {
@@ -7173,7 +7373,7 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) {
})
return
case xerrors.Is(err, errChatProviderNotConfigured):
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
httpapi.Write(ctx, rw, http.StatusPreconditionFailed, codersdk.Response{
Message: "Chat provider is not configured.",
Detail: err.Error(),
})
@@ -7224,15 +7424,40 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) {
}
provider := existing.Provider
if strings.TrimSpace(req.Provider) != "" {
provider = normalizeChatProvider(req.Provider)
if provider == "" {
aiProviderID := existing.AIProviderID
if req.AIProviderID != nil {
//nolint:gocritic // The route already authorized chat model config updates.
aiProvider, err := api.Database.GetAIProviderByID(dbauthz.AsChatd(ctx), *req.AIProviderID)
if err != nil {
if httpapi.Is404Error(err) {
httpapi.Write(ctx, rw, http.StatusPreconditionFailed, codersdk.Response{Message: "AI provider is not configured."})
return
}
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
Message: "Failed to get AI provider.",
Detail: err.Error(),
})
return
}
if !aiProvider.Enabled {
httpapi.Write(ctx, rw, http.StatusPreconditionFailed, codersdk.Response{Message: "AI provider is disabled."})
return
}
provider = string(aiProvider.Type)
aiProviderID = uuid.NullUUID{UUID: aiProvider.ID, Valid: true}
} else if strings.TrimSpace(req.Provider) != "" {
requestedProvider := normalizeChatProvider(req.Provider)
if requestedProvider == "" {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "Invalid provider.",
Detail: chatProviderValidationDetail(),
})
return
}
provider = requestedProvider
if requestedProvider != existing.Provider {
aiProviderID = uuid.NullUUID{}
}
}
model := existing.Model
@@ -7299,14 +7524,30 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) {
ContextLimit: contextLimit,
CompressionThreshold: compressionThreshold,
Options: modelConfigRaw,
AIProviderID: aiProviderID,
UpdatedBy: uuid.NullUUID{UUID: apiKey.UserID, Valid: apiKey.UserID != uuid.Nil},
ID: existing.ID,
}
var updated database.ChatModelConfig
err = api.Database.InTx(func(tx database.Store) error {
if err := requireChatProviderForModelConfig(ctx, tx, updateParams.Provider); err != nil {
return err
if updateParams.AIProviderID.Valid && req.AIProviderID != nil {
//nolint:gocritic // The route already authorized chat model config updates.
aiProvider, err := tx.GetAIProviderByIDForReferenceLock(dbauthz.AsChatd(ctx), updateParams.AIProviderID.UUID)
if err != nil {
if xerrors.Is(err, sql.ErrNoRows) {
return errChatProviderNotConfigured
}
return xerrors.Errorf("get AI provider for update: %w", err)
}
if !aiProvider.Enabled {
return errChatProviderNotConfigured
}
updateParams.Provider = string(aiProvider.Type)
} else if !updateParams.AIProviderID.Valid {
if err := requireChatProviderForModelConfig(ctx, tx, updateParams.Provider); err != nil {
return err
}
}
setAsDefault := updateParams.IsDefault && !existing.IsDefault
@@ -7357,7 +7598,7 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) {
})
return
case xerrors.Is(err, errChatProviderNotConfigured):
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
httpapi.Write(ctx, rw, http.StatusPreconditionFailed, codersdk.Response{
Message: "Chat provider is not configured.",
Detail: err.Error(),
})
@@ -7487,6 +7728,7 @@ func chatModelConfigToUpdateParams(
ContextLimit: config.ContextLimit,
CompressionThreshold: config.CompressionThreshold,
Options: config.Options,
AIProviderID: config.AIProviderID,
UpdatedBy: uuid.NullUUID{},
ID: config.ID,
}
@@ -7589,6 +7831,7 @@ func convertUserChatProviderConfig(
provider database.ChatProvider,
hasUserAPIKey bool,
hasCentralAPIKeyFallback bool,
byokEnabled bool,
) codersdk.UserChatProviderConfig {
displayName := strings.TrimSpace(provider.DisplayName)
if displayName == "" {
@@ -7601,13 +7844,19 @@ func convertUserChatProviderConfig(
DisplayName: displayName,
HasUserAPIKey: hasUserAPIKey,
HasCentralAPIKeyFallback: hasCentralAPIKeyFallback,
BYOKEnabled: byokEnabled,
}
}
func convertChatModelConfig(config database.ChatModelConfig) codersdk.ChatModelConfig {
var aiProviderID *uuid.UUID
if config.AIProviderID.Valid {
aiProviderID = &config.AIProviderID.UUID
}
return codersdk.ChatModelConfig{
ID: config.ID,
Provider: config.Provider,
AIProviderID: aiProviderID,
Model: config.Model,
DisplayName: config.DisplayName,
Enabled: config.Enabled,
@@ -7787,7 +8036,7 @@ const maxChatProviderAPIKeySize = 10240 // 10 KB
func validateChatProviderAPIKeySize(apiKey string) error {
if len(apiKey) > maxChatProviderAPIKeySize {
return xerrors.Errorf("API key exceeds maximum size of %d bytes", maxChatProviderAPIKeySize)
return xerrors.Errorf("API key exceeds maximum size of 10 KB (%d bytes)", maxChatProviderAPIKeySize)
}
return nil
}
+444 -54
View File
@@ -1801,16 +1801,19 @@ func TestListChatModels(t *testing.T) {
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
provider, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "anthropic",
providerType := database.AiProviderTypeAnthropic
chatProvider, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: string(providerType),
CentralAPIKeyEnabled: ptr.Ref(false),
AllowUserAPIKey: ptr.Ref(true),
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, client, string(providerType), "")
contextLimit := int64(4096)
_, err = client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "anthropic",
Provider: string(providerType),
AIProviderID: &aiProvider.ID,
Model: "claude-sonnet",
ContextLimit: &contextLimit,
})
@@ -1821,7 +1824,7 @@ func TestListChatModels(t *testing.T) {
var anthropicProvider *codersdk.ChatModelProvider
for i := range models.Providers {
if models.Providers[i].Provider == "anthropic" {
if models.Providers[i].Provider == string(providerType) {
anthropicProvider = &models.Providers[i]
break
}
@@ -1830,7 +1833,7 @@ func TestListChatModels(t *testing.T) {
require.False(t, anthropicProvider.Available)
require.Equal(t, codersdk.ChatModelProviderUnavailableReasonUserAPIKeyRequired, anthropicProvider.UnavailableReason)
_, err = client.UpsertUserChatProviderKey(ctx, provider.ID, codersdk.CreateUserChatProviderKeyRequest{
_, err = client.UpsertUserChatProviderKey(ctx, chatProvider.ID, codersdk.CreateUserChatProviderKeyRequest{
APIKey: "user-api-key",
})
require.NoError(t, err)
@@ -1856,7 +1859,7 @@ func TestListChatModels(t *testing.T) {
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
provider, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
chatProvider, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "google",
APIKey: "central-api-key",
CentralAPIKeyEnabled: ptr.Ref(true),
@@ -1864,10 +1867,12 @@ func TestListChatModels(t *testing.T) {
AllowCentralAPIKeyFallback: ptr.Ref(true),
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, client, "google", "provider-api-key")
contextLimit := int64(4096)
_, err = client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "google",
AIProviderID: &aiProvider.ID,
Model: "gemini-1.5-pro",
ContextLimit: &contextLimit,
})
@@ -1886,7 +1891,7 @@ func TestListChatModels(t *testing.T) {
require.NotNil(t, googleProvider)
require.True(t, googleProvider.Available)
_, err = client.UpsertUserChatProviderKey(ctx, provider.ID, codersdk.CreateUserChatProviderKeyRequest{
_, err = client.UpsertUserChatProviderKey(ctx, chatProvider.ID, codersdk.CreateUserChatProviderKeyRequest{
APIKey: "user-api-key",
})
require.NoError(t, err)
@@ -1914,15 +1919,17 @@ func TestListChatModels(t *testing.T) {
client := newChatClientWithDeploymentValues(t, values)
_ = coderdtest.CreateFirstUser(t, client.Client)
provider, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
chatProvider, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai",
APIKey: "test-key",
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, client, "openai", "test-key")
contextLimit := int64(4096)
_, err = client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai",
AIProviderID: &aiProvider.ID,
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
})
@@ -1936,7 +1943,7 @@ func TestListChatModels(t *testing.T) {
require.Equal(t, "gpt-4o-mini", models.Providers[0].Models[0].Model)
enabled := false
_, err = client.UpdateChatProvider(ctx, provider.ID, codersdk.UpdateChatProviderConfigRequest{
_, err = client.UpdateChatProvider(ctx, chatProvider.ID, codersdk.UpdateChatProviderConfigRequest{
Enabled: &enabled,
})
require.NoError(t, err)
@@ -2261,6 +2268,186 @@ func TestWatchChats(t *testing.T) {
})
}
func TestUserAIProviderKeys(t *testing.T) {
t.Parallel()
createOpenAIProvider := func(t *testing.T, client *codersdk.ExperimentalClient, name string, enabled bool, apiKeys ...string) codersdk.AIProvider {
t.Helper()
provider, err := client.CreateAIProvider(testutil.Context(t, testutil.WaitLong), codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderTypeOpenAI,
Name: name,
Enabled: enabled,
BaseURL: "https://api.openai.example.com/v1",
APIKeys: apiKeys,
})
require.NoError(t, err)
return provider
}
findUserAIProviderKeyConfig := func(
t *testing.T,
configs []codersdk.UserAIProviderKeyConfig,
providerID uuid.UUID,
) *codersdk.UserAIProviderKeyConfig {
t.Helper()
for i := range configs {
if configs[i].Provider.ID == providerID {
return &configs[i]
}
}
return nil
}
t.Run("SelfServiceLifecycle", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
adminClient := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, adminClient.Client)
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
provider := createOpenAIProvider(t, adminClient, "test-user-key-"+uuid.NewString(), true, "test-provider-api-key")
configs, err := memberClient.ListUserAIProviderKeyConfigs(ctx, "me")
require.NoError(t, err)
cfg := findUserAIProviderKeyConfig(t, configs, provider.ID)
require.NotNil(t, cfg)
require.False(t, cfg.HasUserAPIKey)
require.True(t, cfg.HasProviderAPIKey)
require.True(t, cfg.BYOKEnabled)
cfgValue, err := memberClient.UpsertUserAIProviderKey(ctx, "me", provider.ID, codersdk.CreateUserAIProviderKeyRequest{APIKey: "test-user-api-key"})
require.NoError(t, err)
require.Equal(t, provider.ID, cfgValue.Provider.ID)
require.True(t, cfgValue.HasUserAPIKey)
require.True(t, cfgValue.HasProviderAPIKey)
require.True(t, cfgValue.BYOKEnabled)
configs, err = memberClient.ListUserAIProviderKeyConfigs(ctx, "me")
require.NoError(t, err)
cfg = findUserAIProviderKeyConfig(t, configs, provider.ID)
require.NotNil(t, cfg)
require.True(t, cfg.HasUserAPIKey)
cfgValue, err = memberClient.UpsertUserAIProviderKey(ctx, "me", provider.ID, codersdk.CreateUserAIProviderKeyRequest{APIKey: "replacement-user-api-key"})
require.NoError(t, err)
require.Equal(t, provider.ID, cfgValue.Provider.ID)
require.True(t, cfgValue.HasUserAPIKey)
configs, err = memberClient.ListUserAIProviderKeyConfigs(ctx, "me")
require.NoError(t, err)
cfg = findUserAIProviderKeyConfig(t, configs, provider.ID)
require.NotNil(t, cfg)
require.True(t, cfg.HasUserAPIKey)
require.NoError(t, memberClient.DeleteUserAIProviderKey(ctx, "me", provider.ID))
configs, err = memberClient.ListUserAIProviderKeyConfigs(ctx, "me")
require.NoError(t, err)
cfg = findUserAIProviderKeyConfig(t, configs, provider.ID)
require.NotNil(t, cfg)
require.False(t, cfg.HasUserAPIKey)
})
t.Run("ListsDisabledProviderWithSavedUserKey", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
adminClient := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, adminClient.Client)
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
provider := createOpenAIProvider(t, adminClient, "test-disabled-saved-user-key-"+uuid.NewString(), true)
_, err := memberClient.UpsertUserAIProviderKey(ctx, "me", provider.ID, codersdk.CreateUserAIProviderKeyRequest{APIKey: "test-user-api-key"})
require.NoError(t, err)
enabled := false
_, err = adminClient.UpdateAIProvider(ctx, provider.ID.String(), codersdk.UpdateAIProviderRequest{Enabled: &enabled})
require.NoError(t, err)
configs, err := memberClient.ListUserAIProviderKeyConfigs(ctx, "me")
require.NoError(t, err)
cfg := findUserAIProviderKeyConfig(t, configs, provider.ID)
require.NotNil(t, cfg)
require.False(t, cfg.Provider.Enabled)
require.True(t, cfg.HasUserAPIKey)
})
t.Run("RejectsDisabledProvider", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
adminClient := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, adminClient.Client)
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
provider := createOpenAIProvider(t, adminClient, "test-disabled-user-key-"+uuid.NewString(), false)
_, err := memberClient.UpsertUserAIProviderKey(ctx, "me", provider.ID, codersdk.CreateUserAIProviderKeyRequest{APIKey: "test-user-api-key"})
sdkErr := requireSDKError(t, err, http.StatusPreconditionFailed)
require.Equal(t, "AI provider is disabled.", sdkErr.Message)
})
t.Run("RejectsLargeAPIKey", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
adminClient := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, adminClient.Client)
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
provider := createOpenAIProvider(t, adminClient, "test-large-user-key-"+uuid.NewString(), true)
_, err := memberClient.UpsertUserAIProviderKey(ctx, "me", provider.ID, codersdk.CreateUserAIProviderKeyRequest{APIKey: strings.Repeat("x", 10241)})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "API key too large.", sdkErr.Message)
})
t.Run("RejectsWhitespaceAPIKey", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
adminClient := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, adminClient.Client)
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
provider := createOpenAIProvider(t, adminClient, "test-whitespace-user-key-"+uuid.NewString(), true)
_, err := memberClient.UpsertUserAIProviderKey(ctx, "me", provider.ID, codersdk.CreateUserAIProviderKeyRequest{APIKey: " "})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "API key must not contain leading or trailing whitespace.", sdkErr.Message)
})
t.Run("BYOKDisabledRejectsUpsertAndAllowsDelete", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
values := chatDeploymentValues(t)
values.AI.BridgeConfig.AllowBYOK = serpent.Bool(false)
client := newChatClientWithDeploymentValues(t, values)
_ = coderdtest.CreateFirstUser(t, client.Client)
provider := createOpenAIProvider(t, client, "test-byok-disabled-"+uuid.NewString(), true)
_, err := client.UpsertUserAIProviderKey(ctx, "me", provider.ID, codersdk.CreateUserAIProviderKeyRequest{APIKey: "test-user-api-key"})
sdkErr := requireSDKError(t, err, http.StatusForbidden)
require.Equal(t, "BYOK is disabled.", sdkErr.Message)
configs, err := client.ListUserAIProviderKeyConfigs(ctx, "me")
require.NoError(t, err)
cfg := findUserAIProviderKeyConfig(t, configs, provider.ID)
require.NotNil(t, cfg)
require.False(t, cfg.BYOKEnabled)
require.NoError(t, client.DeleteUserAIProviderKey(ctx, "me", provider.ID))
})
}
func TestListChatProviders(t *testing.T) {
t.Parallel()
@@ -2611,7 +2798,7 @@ func TestCreateChatProvider(t *testing.T) {
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "API key too large.", sdkErr.Message)
require.Equal(t, fmt.Sprintf("API key exceeds maximum size of %d bytes", chatProviderAPIKeySizeLimit), sdkErr.Detail)
require.Equal(t, fmt.Sprintf("API key exceeds maximum size of 10 KB (%d bytes)", chatProviderAPIKeySizeLimit), sdkErr.Detail)
})
t.Run("AllowsMaxSizedAPIKey", func(t *testing.T) {
@@ -2898,7 +3085,7 @@ func TestUpdateChatProvider(t *testing.T) {
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "API key too large.", sdkErr.Message)
require.Equal(t, fmt.Sprintf("API key exceeds maximum size of %d bytes", chatProviderAPIKeySizeLimit), sdkErr.Detail)
require.Equal(t, fmt.Sprintf("API key exceeds maximum size of 10 KB (%d bytes)", chatProviderAPIKeySizeLimit), sdkErr.Detail)
})
t.Run("AllowsMaxSizedAPIKey", func(t *testing.T) {
@@ -2961,11 +3148,13 @@ func TestDeleteChatProvider(t *testing.T) {
AllowUserAPIKey: ptr.Ref(true),
})
require.NoError(t, err)
aiProviderToDelete := createAIProviderForTest(t, client, providerToDelete.Provider, "delete-api-key")
deleteContextLimit := int64(4096)
deleteIsDefault := true
configToDelete, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: providerToDelete.Provider,
AIProviderID: &aiProviderToDelete.ID,
Model: "gpt-4o-delete-provider",
ContextLimit: &deleteContextLimit,
IsDefault: &deleteIsDefault,
@@ -2977,10 +3166,12 @@ func TestDeleteChatProvider(t *testing.T) {
APIKey: "keep-api-key",
})
require.NoError(t, err)
keepAIProvider := createAIProviderForTest(t, client, keepProvider.Provider, "keep-api-key")
keepContextLimit := int64(8192)
keepConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: keepProvider.Provider,
AIProviderID: &keepAIProvider.ID,
Model: "claude-keep-provider",
ContextLimit: &keepContextLimit,
})
@@ -3082,11 +3273,13 @@ func TestDeleteChatProvider(t *testing.T) {
APIKey: "only-provider-api-key",
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, client, provider.Provider, "only-provider-api-key")
contextLimit := int64(4096)
isDefault := true
config, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: provider.Provider,
AIProviderID: &aiProvider.ID,
Model: "gpt-4o-only-provider",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
@@ -3636,7 +3829,7 @@ func TestUpsertUserChatProviderKey(t *testing.T) {
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "API key too large.", sdkErr.Message)
require.Equal(t, fmt.Sprintf("API key exceeds maximum size of %d bytes", chatProviderAPIKeySizeLimit), sdkErr.Detail)
require.Equal(t, fmt.Sprintf("API key exceeds maximum size of 10 KB (%d bytes)", chatProviderAPIKeySizeLimit), sdkErr.Detail)
})
t.Run("AllowsMaxSizedAPIKey", func(t *testing.T) {
@@ -3695,16 +3888,13 @@ func TestListChatModelConfigs(t *testing.T) {
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
_, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai",
APIKey: "test-api-key",
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, client, "openai", "test-api-key")
contextLimit := int64(4096)
enabled := false
disabledConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai",
AIProviderID: &aiProvider.ID,
Model: "gpt-4o-disabled",
DisplayName: "GPT-4o Disabled",
Enabled: &enabled,
@@ -3741,6 +3931,7 @@ func TestListChatModelConfigs(t *testing.T) {
enabled := false
_, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: enabledConfig.Provider,
AIProviderID: enabledConfig.AIProviderID,
Model: "gpt-4o-disabled",
DisplayName: "GPT-4o Disabled",
Enabled: &enabled,
@@ -3762,15 +3953,12 @@ func TestListChatModelConfigs(t *testing.T) {
client, db := newChatClientWithDatabase(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
_, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai",
APIKey: "test-api-key",
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, client, "openai", "test-api-key")
legacyOptions := json.RawMessage(`{"input_price_per_million_tokens":0.15,"output_price_per_million_tokens":0.6,"cache_read_price_per_million_tokens":0.03,"cache_write_price_per_million_tokens":0.3}`)
storedConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
Provider: "openai",
AIProviderID: uuid.NullUUID{UUID: aiProvider.ID, Valid: true},
Model: "gpt-4o-mini-legacy",
DisplayName: "GPT-4o Mini Legacy",
CreatedBy: uuid.NullUUID{UUID: firstUser.UserID, Valid: true},
@@ -3831,11 +4019,7 @@ func TestCreateChatModelConfig(t *testing.T) {
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
_, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai",
APIKey: "test-api-key",
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, client, "openai", "test-api-key")
contextLimit := int64(4096)
isDefault := true
@@ -3849,6 +4033,7 @@ func TestCreateChatModelConfig(t *testing.T) {
}
modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai",
AIProviderID: &aiProvider.ID,
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
@@ -3875,15 +4060,12 @@ func TestCreateChatModelConfig(t *testing.T) {
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
_, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai",
APIKey: "test-api-key",
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, client, "openai", "test-api-key")
contextLimit := int64(4096)
_, err = client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai",
AIProviderID: &aiProvider.ID,
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
ModelConfig: &codersdk.ChatModelCallConfig{
@@ -3907,16 +4089,18 @@ func TestCreateChatModelConfig(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
aiProvider := createAIProviderForTest(t, client, "openai", "test-api-key")
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai",
Model: "gpt-4o-mini",
Provider: "openai",
AIProviderID: &aiProvider.ID,
Model: "gpt-4o-mini",
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "Context limit is required.", sdkErr.Message)
})
t.Run("ProviderNotConfigured", func(t *testing.T) {
t.Run("AIProviderIDRequired", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
@@ -3930,7 +4114,94 @@ func TestCreateChatModelConfig(t *testing.T) {
ContextLimit: &contextLimit,
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "Chat provider is not configured.", sdkErr.Message)
require.Equal(t, "AI provider ID is required.", sdkErr.Message)
})
t.Run("ProviderNotConfigured", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
contextLimit := int64(4096)
missingProviderID := uuid.New()
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai",
AIProviderID: &missingProviderID,
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
})
sdkErr := requireSDKError(t, err, http.StatusPreconditionFailed)
require.Equal(t, "AI provider is not configured.", sdkErr.Message)
})
t.Run("WithAIProviderID", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
provider, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderTypeOpenAI,
Name: "test-model-config-provider-" + uuid.NewString(),
Enabled: true,
BaseURL: "https://api.openai.com/v1",
})
require.NoError(t, err)
contextLimit := int64(4096)
modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
AIProviderID: &provider.ID,
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
})
require.NoError(t, err)
require.Equal(t, "openai", modelConfig.Provider)
require.NotNil(t, modelConfig.AIProviderID)
require.Equal(t, provider.ID, *modelConfig.AIProviderID)
})
t.Run("AIProviderIDNotConfigured", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
missingProviderID := uuid.New()
contextLimit := int64(4096)
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
AIProviderID: &missingProviderID,
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
})
sdkErr := requireSDKError(t, err, http.StatusPreconditionFailed)
require.Equal(t, "AI provider is not configured.", sdkErr.Message)
})
t.Run("AIProviderIDDisabled", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
provider, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderTypeOpenAI,
Name: "test-disabled-model-provider-" + uuid.NewString(),
Enabled: false,
BaseURL: "https://api.openai.com/v1",
})
require.NoError(t, err)
contextLimit := int64(4096)
_, err = client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
AIProviderID: &provider.ID,
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
})
sdkErr := requireSDKError(t, err, http.StatusPreconditionFailed)
require.Equal(t, "AI provider is disabled.", sdkErr.Message)
})
t.Run("ForbiddenForOrganizationMember", func(t *testing.T) {
@@ -3942,15 +4213,12 @@ func TestCreateChatModelConfig(t *testing.T) {
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
_, err := adminClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai",
APIKey: "test-api-key",
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, adminClient, "openai", "test-api-key")
contextLimit := int64(4096)
_, err = memberClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
_, err := memberClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai",
AIProviderID: &aiProvider.ID,
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
})
@@ -4041,16 +4309,13 @@ func TestUpdateChatModelConfig(t *testing.T) {
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
_, err := adminClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai",
APIKey: "test-api-key",
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, adminClient, "openai", "test-api-key")
contextLimit := int64(4096)
enabled := false
modelConfig, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai",
AIProviderID: &aiProvider.ID,
Model: "gpt-4o-reenable",
DisplayName: "GPT-4o Re-enable",
Enabled: &enabled,
@@ -4115,6 +4380,100 @@ func TestUpdateChatModelConfig(t *testing.T) {
)
})
t.Run("UpdateAIProviderID", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
provider, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderTypeAnthropic,
Name: "test-update-model-provider-" + uuid.NewString(),
Enabled: true,
BaseURL: "https://api.anthropic.com",
})
require.NoError(t, err)
updated, err := client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
AIProviderID: &provider.ID,
Model: "claude-3-5-sonnet-latest",
})
require.NoError(t, err)
require.Equal(t, "anthropic", updated.Provider)
require.NotNil(t, updated.AIProviderID)
require.Equal(t, provider.ID, *updated.AIProviderID)
})
t.Run("UpdateProviderPreservesAIProviderIDWhenTypeUnchanged", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
provider, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderTypeAnthropic,
Name: "test-preserve-model-provider-" + uuid.NewString(),
Enabled: true,
BaseURL: "https://api.anthropic.com",
})
require.NoError(t, err)
updated, err := client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
AIProviderID: &provider.ID,
Model: "claude-3-5-sonnet-latest",
})
require.NoError(t, err)
require.NotNil(t, updated.AIProviderID)
updated, err = client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
Provider: "anthropic",
Model: "claude-3-5-haiku-latest",
})
require.NoError(t, err)
require.NotNil(t, updated.AIProviderID)
require.Equal(t, provider.ID, *updated.AIProviderID)
})
t.Run("UpdateAIProviderIDNotConfigured", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
missingProviderID := uuid.New()
_, err := client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
AIProviderID: &missingProviderID,
})
sdkErr := requireSDKError(t, err, http.StatusPreconditionFailed)
require.Equal(t, "AI provider is not configured.", sdkErr.Message)
})
t.Run("UpdateAIProviderIDDisabled", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
modelConfig := createChatModelConfig(t, client)
provider, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderTypeOpenAI,
Name: "test-update-disabled-model-provider-" + uuid.NewString(),
Enabled: false,
BaseURL: "https://api.openai.com/v1",
})
require.NoError(t, err)
_, err = client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
AIProviderID: &provider.ID,
})
sdkErr := requireSDKError(t, err, http.StatusPreconditionFailed)
require.Equal(t, "AI provider is disabled.", sdkErr.Message)
})
t.Run("ProviderNotConfigured", func(t *testing.T) {
t.Parallel()
@@ -4126,7 +4485,7 @@ func TestUpdateChatModelConfig(t *testing.T) {
_, err := client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
Provider: "anthropic",
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
sdkErr := requireSDKError(t, err, http.StatusPreconditionFailed)
require.Equal(t, "Chat provider is not configured.", sdkErr.Message)
})
@@ -4167,16 +4526,13 @@ func TestUpdateChatModelConfig(t *testing.T) {
_ = coderdtest.CreateFirstUser(t, client.Client)
defaultConfig := createChatModelConfig(t, client)
_, err := client.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "anthropic",
APIKey: "candidate-api-key",
})
require.NoError(t, err)
aiProvider := createAIProviderForTest(t, client, "anthropic", "candidate-api-key")
contextLimit := int64(4096)
isDefault := false
candidateConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "anthropic",
AIProviderID: &aiProvider.ID,
Model: "claude-3-5-sonnet",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
@@ -10495,6 +10851,38 @@ func TestWatchChatGitAuthz(t *testing.T) {
require.Equal(t, http.StatusForbidden, res.StatusCode)
}
func createAIProviderForTest(
t testing.TB,
client *codersdk.ExperimentalClient,
provider string,
apiKey string,
) codersdk.AIProvider {
t.Helper()
ctx := testutil.Context(t, testutil.WaitLong)
req := codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderType(provider),
Name: "test-" + provider + "-" + uuid.NewString(),
BaseURL: aiProviderBaseURLForTest(provider),
Enabled: true,
}
if apiKey != "" {
req.APIKeys = []string{apiKey}
}
aiProvider, err := client.CreateAIProvider(ctx, req)
require.NoError(t, err)
return aiProvider
}
func aiProviderBaseURLForTest(provider string) string {
switch provider {
case "anthropic", "bedrock", "google":
return "https://api.example.com"
default:
return "https://api.example.com/v1"
}
}
func createChatModelConfig(t testing.TB, client *codersdk.ExperimentalClient) codersdk.ChatModelConfig {
t.Helper()
return coderdtest.CreateOpenAICompatChatModelConfig(t, client, "")
@@ -10529,10 +10917,12 @@ func createAdditionalChatModelConfig(
t.Helper()
ctx := testutil.Context(t, testutil.WaitLong)
aiProvider := createAIProviderForTest(t, client, provider, "test-api-key")
contextLimit := int64(4096)
isDefault := false
modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: provider,
AIProviderID: &aiProvider.ID,
Model: model,
ContextLimit: &contextLimit,
IsDefault: &isDefault,
+1
View File
@@ -84,6 +84,7 @@ const (
SubjectTypeBoundaryUsageTracker SubjectType = "boundary_usage_tracker"
SubjectTypeWorkspaceBuilder SubjectType = "workspace_builder"
SubjectTypeChatd SubjectType = "chatd"
SubjectTypeAIProviderMetadataReader SubjectType = "ai_provider_metadata_reader"
)
const (
+7 -7
View File
@@ -350,13 +350,13 @@ func (p *Server) resolveAdvisorModelOverride(
return fallbackModel, fallbackCallConfig
}
// GetEnabledChatModelConfigByID joins on chat_providers.enabled = TRUE
// and chat_model_configs.enabled = TRUE, so it returns sql.ErrNoRows
// the moment an admin disables either the model config or its provider.
// Using the cached ModelConfigByID here would keep resolving an override
// whose provider was just disabled, and an env or central fallback key
// would let ModelFromConfig succeed, silently routing advisor prompts
// to a provider the admin expects to be off.
// GetEnabledChatModelConfigByID checks the model config and referenced
// provider enabled state, so it returns sql.ErrNoRows the moment an
// admin disables either one. Using the cached ModelConfigByID here
// would keep resolving an override whose provider was just disabled,
// and an available fallback key would let ModelFromConfig succeed,
// silently routing advisor prompts to a provider the admin expects to
// be off.
overrideConfig, err := p.db.GetEnabledChatModelConfigByID(
ctx,
advisorCfg.ModelConfigID,
+7 -97
View File
@@ -288,22 +288,7 @@ func TestSubagentChatExcludesWorkspaceProvisioningTools(t *testing.T) {
)
})
_, err := expClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai-compat",
APIKey: "test-api-key",
BaseURL: openAIURL,
})
require.NoError(t, err)
contextLimit := int64(4096)
isDefault := true
_, err = expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai-compat",
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
})
require.NoError(t, err)
coderdtest.CreateOpenAICompatChatModelConfig(t, expClient, openAIURL)
// Create a root chat whose first model call will spawn a subagent.
chat, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
@@ -483,22 +468,7 @@ func TestPlanModeSubagentChatExcludesAskUserQuestion(t *testing.T) {
)
})
_, err = expClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai-compat",
APIKey: "test-api-key",
BaseURL: openAIURL,
})
require.NoError(t, err)
contextLimit := int64(4096)
isDefault := true
_, err = expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai-compat",
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
})
require.NoError(t, err)
coderdtest.CreateOpenAICompatChatModelConfig(t, expClient, openAIURL)
chat, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: user.OrganizationID,
@@ -638,24 +608,9 @@ func TestExploreSubagentIsReadOnly(t *testing.T) {
)
})
_, err := expClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai-compat",
APIKey: "test-api-key",
BaseURL: openAIURL,
})
require.NoError(t, err)
coderdtest.CreateOpenAICompatChatModelConfig(t, expClient, openAIURL)
contextLimit := int64(4096)
isDefault := true
_, err = expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai-compat",
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
})
require.NoError(t, err)
_, err = expClient.CreateChat(ctx, codersdk.CreateChatRequest{
_, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: user.OrganizationID,
WorkspaceID: &workspace.ID,
Content: []codersdk.ChatInputPart{
@@ -4953,22 +4908,7 @@ func TestCreateWorkspaceTool_EndToEnd(t *testing.T) {
)
})
_, err := expClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai-compat",
APIKey: "test-api-key",
BaseURL: openAIURL,
})
require.NoError(t, err)
contextLimit := int64(4096)
isDefault := true
_, err = expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai-compat",
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
})
require.NoError(t, err)
coderdtest.CreateOpenAICompatChatModelConfig(t, expClient, openAIURL)
chat, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: user.OrganizationID,
@@ -5123,22 +5063,7 @@ func TestStartWorkspaceTool_EndToEnd(t *testing.T) {
)
})
_, err := expClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai-compat",
APIKey: "test-api-key",
BaseURL: openAIURL,
})
require.NoError(t, err)
contextLimit := int64(4096)
isDefault := true
_, err = expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai-compat",
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
})
require.NoError(t, err)
coderdtest.CreateOpenAICompatChatModelConfig(t, expClient, openAIURL)
// Create a chat with the stopped workspace pre-associated.
chat, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
@@ -8586,22 +8511,7 @@ func TestAgentContextFilesAndSkillsLoadedIntoChat(t *testing.T) {
)
})
_, err := expClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai-compat",
APIKey: "test-api-key",
BaseURL: openAIURL,
})
require.NoError(t, err)
contextLimit := int64(4096)
isDefault := true
_, err = expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai-compat",
Model: "gpt-4o-mini",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
})
require.NoError(t, err)
coderdtest.CreateOpenAICompatChatModelConfig(t, expClient, openAIURL)
workspaceID := workspace.ID
chat, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
+65 -27
View File
@@ -5,6 +5,7 @@ import (
"os"
"testing"
"github.com/google/uuid"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/coderd/coderdtest"
@@ -13,6 +14,52 @@ import (
"github.com/coder/coder/v2/testutil"
)
func createIntegrationAIProvider(
ctx context.Context,
t testing.TB,
client *codersdk.ExperimentalClient,
providerType codersdk.AIProviderType,
apiKey string,
baseURL string,
) codersdk.AIProvider {
t.Helper()
if baseURL == "" {
baseURL = defaultIntegrationAIProviderBaseURL(providerType)
}
provider, err := client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{
Type: providerType,
Name: string(providerType) + "-" + uuid.NewString(),
DisplayName: aiProviderDisplayName(providerType),
Enabled: true,
BaseURL: baseURL,
APIKeys: []string{apiKey},
})
require.NoError(t, err)
return provider
}
func defaultIntegrationAIProviderBaseURL(providerType codersdk.AIProviderType) string {
switch providerType {
case codersdk.AIProviderTypeAnthropic:
return "https://api.anthropic.com"
case codersdk.AIProviderTypeOpenAI:
return "https://api.openai.com/v1"
default:
return "https://api.example.com"
}
}
func aiProviderDisplayName(providerType codersdk.AIProviderType) string {
switch providerType {
case codersdk.AIProviderTypeAnthropic:
return "Anthropic"
case codersdk.AIProviderTypeOpenAI:
return "OpenAI"
default:
return string(providerType)
}
}
// TestAnthropicWebSearchRoundTrip is an integration test that verifies
// provider-executed tool results (web_search) survive the full
// persist → reconstruct → re-send cycle. It sends a query that
@@ -43,19 +90,16 @@ func TestAnthropicWebSearchRoundTrip(t *testing.T) {
user := coderdtest.CreateFirstUser(t, client)
expClient := codersdk.NewExperimentalClient(client)
// Configure an Anthropic provider with the real API key.
_, err := expClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "anthropic",
APIKey: apiKey,
BaseURL: baseURL,
})
require.NoError(t, err)
provider := createIntegrationAIProvider(
ctx, t, expClient, codersdk.AIProviderTypeAnthropic, apiKey, baseURL,
)
// Create a model config that enables web_search.
contextLimit := int64(200000)
isDefault := true
_, err = expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "anthropic",
_, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: string(provider.Type),
AIProviderID: &provider.ID,
Model: "claude-sonnet-4-20250514",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
@@ -303,13 +347,9 @@ func TestOpenAIReasoningRoundTrip(t *testing.T) {
user := coderdtest.CreateFirstUser(t, client)
expClient := codersdk.NewExperimentalClient(client)
// Configure an OpenAI provider with the real API key.
_, err := expClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai",
APIKey: apiKey,
BaseURL: baseURL,
})
require.NoError(t, err)
provider := createIntegrationAIProvider(
ctx, t, expClient, codersdk.AIProviderTypeOpenAI, apiKey, baseURL,
)
// Create a model config for a reasoning model with Store: true
// (the default). Using o4-mini because it always produces
@@ -317,8 +357,9 @@ func TestOpenAIReasoningRoundTrip(t *testing.T) {
contextLimit := int64(200000)
isDefault := true
reasoningSummary := "auto"
_, err = expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai",
_, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: string(provider.Type),
AIProviderID: &provider.ID,
Model: "o4-mini",
ContextLimit: &contextLimit,
IsDefault: &isDefault,
@@ -457,21 +498,18 @@ func TestOpenAIReasoningRoundTripStoreFalse(t *testing.T) {
user := coderdtest.CreateFirstUser(t, client)
expClient := codersdk.NewExperimentalClient(client)
// Configure an OpenAI provider with the real API key.
_, err := expClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
Provider: "openai",
APIKey: apiKey,
BaseURL: baseURL,
})
require.NoError(t, err)
provider := createIntegrationAIProvider(
ctx, t, expClient, codersdk.AIProviderTypeOpenAI, apiKey, baseURL,
)
// Create a model config for a reasoning model with Store: false.
// Using o4-mini because it always produces reasoning items.
contextLimit := int64(200000)
isDefault := true
reasoningSummary := "auto"
_, err = expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: "openai",
_, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
Provider: string(provider.Type),
AIProviderID: &provider.ID,
Model: "o4-mini",
ContextLimit: &contextLimit,
IsDefault: &isDefault,