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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user