mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: use AI provider chat APIs (#25415)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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(¶ms)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
*
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -84,6 +84,7 @@ const (
|
||||
SubjectTypeBoundaryUsageTracker SubjectType = "boundary_usage_tracker"
|
||||
SubjectTypeWorkspaceBuilder SubjectType = "workspace_builder"
|
||||
SubjectTypeChatd SubjectType = "chatd"
|
||||
SubjectTypeAIProviderMetadataReader SubjectType = "ai_provider_metadata_reader"
|
||||
)
|
||||
|
||||
const (
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user