mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
refactor: drop chat_model_configs provider column (#26877)
The provider type already lives authoritatively in ai_providers.type, reachable on every active row through ai_provider_id, which the chat_model_configs_ai_provider_required_when_active CHECK makes mandatory. The stored provider string was a denormalized copy the system kept in sync with a startup backfill and no longer needs. Every surface now derives provider type from the linked ai_providers row. Telemetry is the one exception: it keeps emitting provider, now sourced from ai_providers.type via a JOIN, so the BigQuery column and the Nexus dashboards that read it are unaffected. The experimental HTTP/SDK response drops provider and makes ai_provider_id required, since those endpoints return only active configs; consumers resolve provider type from ai_provider_id and the AI providers listing. This ships in a single release with no compatibility window: production reads the table via SELECT *, so a pre-drop binary fails config reads the moment the column is gone. Operators must scale to zero before upgrading, and there is no rollback. Closes CODAGT-599
This commit is contained in:
@@ -9,16 +9,12 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
// BackfillBedrockProviderType promotes legacy ai_providers rows stored as
|
||||
// type=anthropic with Bedrock settings to type=bedrock. Must run after newAPI
|
||||
// so options.Database is dbcrypt-wrapped. Idempotent; errors are logged and
|
||||
// startup continues.
|
||||
//
|
||||
// BackfillChatModelConfigProviderStrings must run after this function so
|
||||
// provider types are correct when its JOIN executes.
|
||||
func BackfillBedrockProviderType(ctx context.Context, db database.Store, logger slog.Logger) {
|
||||
//nolint:gocritic // Startup-only backfill; no user actor is present.
|
||||
sysCtx := dbauthz.AsSystemRestricted(ctx)
|
||||
@@ -70,25 +66,3 @@ func BackfillBedrockProviderType(ctx context.Context, db database.Store, logger
|
||||
logger.Info(ctx, "backfilled bedrock provider types", slog.F("count", promoted))
|
||||
}
|
||||
}
|
||||
|
||||
// BackfillChatModelConfigProviderStrings fixes stale chat_model_configs.provider
|
||||
// strings left as "anthropic" when the linked provider was promoted from
|
||||
// type=anthropic to type=bedrock by BackfillBedrockProviderType. Errors are
|
||||
// logged and startup continues.
|
||||
func BackfillChatModelConfigProviderStrings(ctx context.Context, db database.Store, logger slog.Logger) {
|
||||
//nolint:gocritic // Startup-only backfill; no user actor is present.
|
||||
sysCtx := dbauthz.AsSystemRestricted(ctx)
|
||||
result, err := db.BackfillChatModelConfigProvider(sysCtx, database.BackfillChatModelConfigProviderParams{
|
||||
OldProvider: string(codersdk.AIProviderTypeAnthropic),
|
||||
NewProvider: string(codersdk.AIProviderTypeBedrock),
|
||||
})
|
||||
if err != nil {
|
||||
logger.Error(ctx, "backfill chat model config provider strings", slog.Error(err))
|
||||
return
|
||||
}
|
||||
if result != nil {
|
||||
if n, _ := result.RowsAffected(); n > 0 {
|
||||
logger.Info(ctx, "backfilled chat model config provider strings", slog.F("count", n))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
@@ -193,112 +192,6 @@ func TestBackfillBedrockProviderType(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, database.AIProviderTypeBedrock, goodRow.Type, "valid row alongside unparsable one must still be promoted")
|
||||
})
|
||||
|
||||
// --- chat_model_configs.provider backfill ---
|
||||
// These subtests rely on the DB already having type=bedrock providers
|
||||
// from the provider backfill subtests above.
|
||||
|
||||
t.Run("FixesStaleModelConfigProvider", func(t *testing.T) {
|
||||
// Simulate a model config created when the linked provider was still
|
||||
// type=anthropic. The stored provider string is "anthropic" but the
|
||||
// linked provider row now has type=bedrock.
|
||||
bedrockProvider := dbgen.AIProvider(t, db, database.AIProvider{
|
||||
Type: database.AIProviderTypeBedrock,
|
||||
Settings: bedrockSettings,
|
||||
})
|
||||
staleConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: uuid.NullUUID{UUID: bedrockProvider.ID, Valid: true},
|
||||
})
|
||||
|
||||
coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger)
|
||||
|
||||
updated, err := db.GetChatModelConfigByID(ctx, staleConfig.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "bedrock", updated.Provider, "stale anthropic provider string must be fixed to bedrock")
|
||||
|
||||
// Second run must be a no-op: the same config must still be "bedrock".
|
||||
coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger)
|
||||
|
||||
updated, err = db.GetChatModelConfigByID(ctx, staleConfig.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "bedrock", updated.Provider, "provider must remain bedrock after second run")
|
||||
})
|
||||
|
||||
t.Run("ModelConfigIdempotent", func(t *testing.T) {
|
||||
before, err := db.GetChatModelConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger)
|
||||
|
||||
after, err := db.GetChatModelConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, len(before), len(after), "second run must not create or delete rows")
|
||||
})
|
||||
|
||||
t.Run("PreservesNonAnthropicModelConfig", func(t *testing.T) {
|
||||
// A model config with provider="openai" linked to a Bedrock provider
|
||||
// must not be touched. Only "anthropic" → "bedrock" is in scope.
|
||||
bedrockProvider := dbgen.AIProvider(t, db, database.AIProvider{
|
||||
Type: database.AIProviderTypeBedrock,
|
||||
Settings: bedrockSettings,
|
||||
})
|
||||
openAIConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
AIProviderID: uuid.NullUUID{UUID: bedrockProvider.ID, Valid: true},
|
||||
})
|
||||
|
||||
coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger)
|
||||
|
||||
row, err := db.GetChatModelConfigByID(ctx, openAIConfig.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "openai", row.Provider, "non-anthropic provider string must not be changed")
|
||||
})
|
||||
|
||||
t.Run("SkipsModelConfigWithDeletedProvider", func(t *testing.T) {
|
||||
// Verifies the EXISTS subquery excludes soft-deleted providers.
|
||||
// The model config provider string must stay "anthropic" because
|
||||
// the linked provider is deleted and therefore excluded by the
|
||||
// AND deleted = FALSE condition in the query.
|
||||
deletedProvider := dbgen.AIProvider(t, db, database.AIProvider{
|
||||
Type: database.AIProviderTypeBedrock,
|
||||
Settings: bedrockSettings,
|
||||
})
|
||||
staleConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: uuid.NullUUID{UUID: deletedProvider.ID, Valid: true},
|
||||
})
|
||||
require.NoError(t, db.DeleteAIProviderByID(ctx, deletedProvider.ID))
|
||||
|
||||
coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger)
|
||||
|
||||
row, err := db.GetChatModelConfigByID(ctx, staleConfig.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "anthropic", row.Provider, "config linked to deleted provider must not be updated")
|
||||
})
|
||||
|
||||
t.Run("SkipsDeletedModelConfig", func(t *testing.T) {
|
||||
// The SQL query guards on deleted = FALSE. Capture the config ID
|
||||
// before deletion so we delete the right row regardless of ordering.
|
||||
bedrockProvider := dbgen.AIProvider(t, db, database.AIProvider{
|
||||
Type: database.AIProviderTypeBedrock,
|
||||
Settings: bedrockSettings,
|
||||
})
|
||||
cfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: uuid.NullUUID{UUID: bedrockProvider.ID, Valid: true},
|
||||
})
|
||||
|
||||
before, err := db.GetChatModelConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, db.DeleteChatModelConfigByID(ctx, cfg.ID))
|
||||
|
||||
coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger)
|
||||
|
||||
after, err := db.GetChatModelConfigs(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, len(before)-1, len(after), "deleted config must not reappear after backfill")
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("ListFailure", func(t *testing.T) {
|
||||
@@ -352,17 +245,4 @@ func TestBackfillBedrockProviderType(t *testing.T) {
|
||||
// ErrNoRows is benign: provider was deleted between list and update.
|
||||
coderd.BackfillBedrockProviderType(ctx, db, testLogger(t))
|
||||
})
|
||||
|
||||
t.Run("ModelConfigQueryFailure", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
|
||||
db.EXPECT().
|
||||
BackfillChatModelConfigProvider(gomock.Any(), gomock.Any()).
|
||||
Return(nil, sql.ErrConnDone)
|
||||
|
||||
coderd.BackfillChatModelConfigProviderStrings(ctx, db, testLogger(t))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -73,7 +73,6 @@ func CreateOpenAICompatChatModelConfig(
|
||||
contextLimit := int64(4096)
|
||||
isDefault := true
|
||||
modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: TestChatProviderOpenAICompat,
|
||||
AIProviderID: &provider.ID,
|
||||
Model: TestChatModelOpenAICompat,
|
||||
ContextLimit: &contextLimit,
|
||||
|
||||
@@ -1724,13 +1724,6 @@ func (q *querier) AutoArchiveInactiveChats(ctx context.Context, arg database.Aut
|
||||
return q.db.AutoArchiveInactiveChats(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) BackfillChatModelConfigProvider(ctx context.Context, arg database.BackfillChatModelConfigProviderParams) (sql.Result, error) {
|
||||
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q.db.BackfillChatModelConfigProvider(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) BackoffChatDiffStatus(ctx context.Context, arg database.BackoffChatDiffStatusParams) error {
|
||||
// This is a system-level operation used by the gitsync
|
||||
// background worker to reschedule failed refreshes. Same
|
||||
@@ -2113,13 +2106,6 @@ func (q *querier) DeleteChatModelConfigsByAIProviderID(ctx context.Context, aiPr
|
||||
return q.db.DeleteChatModelConfigsByAIProviderID(ctx, aiProviderID)
|
||||
}
|
||||
|
||||
func (q *querier) DeleteChatModelConfigsByProvider(ctx context.Context, provider string) error {
|
||||
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
|
||||
return err
|
||||
}
|
||||
return q.db.DeleteChatModelConfigsByProvider(ctx, provider)
|
||||
}
|
||||
|
||||
func (q *querier) DeleteChatQueuedMessage(ctx context.Context, arg database.DeleteChatQueuedMessageParams) error {
|
||||
chat, err := q.db.GetChatByID(ctx, arg.ChatID)
|
||||
if err != nil {
|
||||
@@ -3677,7 +3663,7 @@ func (q *querier) GetEnabledChatModelConfigByID(ctx context.Context, id uuid.UUI
|
||||
return q.db.GetEnabledChatModelConfigByID(ctx, id)
|
||||
}
|
||||
|
||||
func (q *querier) GetEnabledChatModelConfigs(ctx context.Context) ([]database.ChatModelConfig, error) {
|
||||
func (q *querier) GetEnabledChatModelConfigs(ctx context.Context) ([]database.GetEnabledChatModelConfigsRow, error) {
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -657,11 +657,6 @@ func (s *MethodTestSuite) TestChats() {
|
||||
dbm.EXPECT().DeleteChatModelConfigByID(gomock.Any(), id).Return(nil).AnyTimes()
|
||||
check.Args(id).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate)
|
||||
}))
|
||||
s.Run("DeleteChatModelConfigsByProvider", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
providerName := "test-provider"
|
||||
dbm.EXPECT().DeleteChatModelConfigsByProvider(gomock.Any(), providerName).Return(nil).AnyTimes()
|
||||
check.Args(providerName).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate)
|
||||
}))
|
||||
s.Run("DeleteChatModelConfigsByAIProviderID", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
providerID := uuid.New()
|
||||
dbm.EXPECT().DeleteChatModelConfigsByAIProviderID(gomock.Any(), providerID).Return(nil).AnyTimes()
|
||||
@@ -1212,10 +1207,10 @@ func (s *MethodTestSuite) TestChats() {
|
||||
check.Args(config.ID).Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(config)
|
||||
}))
|
||||
s.Run("GetEnabledChatModelConfigs", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
configA := testutil.Fake(s.T(), faker, database.ChatModelConfig{})
|
||||
configB := testutil.Fake(s.T(), faker, database.ChatModelConfig{})
|
||||
dbm.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.ChatModelConfig{configA, configB}, nil).AnyTimes()
|
||||
check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns([]database.ChatModelConfig{configA, configB})
|
||||
rowA := testutil.Fake(s.T(), faker, database.GetEnabledChatModelConfigsRow{})
|
||||
rowB := testutil.Fake(s.T(), faker, database.GetEnabledChatModelConfigsRow{})
|
||||
dbm.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.GetEnabledChatModelConfigsRow{rowA, rowB}, nil).AnyTimes()
|
||||
check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns([]database.GetEnabledChatModelConfigsRow{rowA, rowB})
|
||||
}))
|
||||
|
||||
s.Run("GetStaleChats", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
@@ -1257,12 +1252,11 @@ func (s *MethodTestSuite) TestChats() {
|
||||
}))
|
||||
s.Run("InsertChatModelConfig", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
arg := database.InsertChatModelConfigParams{
|
||||
Provider: "test-provider",
|
||||
Model: "test-model",
|
||||
DisplayName: "Test Model",
|
||||
Enabled: true,
|
||||
}
|
||||
config := testutil.Fake(s.T(), faker, database.ChatModelConfig{Provider: arg.Provider, Model: arg.Model, DisplayName: arg.DisplayName, Enabled: arg.Enabled})
|
||||
config := testutil.Fake(s.T(), faker, database.ChatModelConfig{Model: arg.Model, DisplayName: arg.DisplayName, Enabled: arg.Enabled})
|
||||
dbm.EXPECT().InsertChatModelConfig(gomock.Any(), arg).Return(config, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(config)
|
||||
}))
|
||||
@@ -1515,7 +1509,6 @@ func (s *MethodTestSuite) TestChats() {
|
||||
config := testutil.Fake(s.T(), faker, database.ChatModelConfig{})
|
||||
arg := database.UpdateChatModelConfigParams{
|
||||
ID: config.ID,
|
||||
Provider: "updated-provider",
|
||||
Model: "updated-model",
|
||||
DisplayName: "Updated Model",
|
||||
Enabled: true,
|
||||
@@ -6859,14 +6852,6 @@ func (s *MethodTestSuite) TestAIBridge() {
|
||||
dbm.EXPECT().DeleteAIProviderByID(gomock.Any(), provider.ID).Return(nil).AnyTimes()
|
||||
check.Args(provider.ID).Asserts(rbac.ResourceAIProvider, policy.ActionDelete).Returns()
|
||||
}))
|
||||
s.Run("BackfillChatModelConfigProvider", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
arg := database.BackfillChatModelConfigProviderParams{
|
||||
OldProvider: "anthropic",
|
||||
NewProvider: "bedrock",
|
||||
}
|
||||
dbm.EXPECT().BackfillChatModelConfigProvider(gomock.Any(), arg).Return(nil, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate)
|
||||
}))
|
||||
s.Run("UpdateEncryptedAIProviderSettings", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
provider := testutil.Fake(s.T(), faker, database.AIProvider{})
|
||||
arg := database.UpdateEncryptedAIProviderSettingsParams{
|
||||
|
||||
@@ -160,14 +160,15 @@ const (
|
||||
|
||||
func ChatModelConfig(t testing.TB, db database.Store, seed database.ChatModelConfig, munge ...func(*database.InsertChatModelConfigParams)) database.ChatModelConfig {
|
||||
t.Helper()
|
||||
providerName := takeFirst(seed.Provider, "openai")
|
||||
aiProviderID := seed.AIProviderID
|
||||
if !aiProviderID.Valid {
|
||||
// No AIProviderID supplied: reuse or create a default openai provider.
|
||||
// Tests needing a specific provider type should pass seed.AIProviderID.
|
||||
providers, err := db.GetAIProviders(genCtx, database.GetAIProvidersParams{IncludeDisabled: true})
|
||||
require.NoError(t, err, "get ai providers")
|
||||
var provider database.AIProvider
|
||||
for _, candidate := range providers {
|
||||
if candidate.Type != database.AIProviderType(providerName) {
|
||||
if candidate.Type != database.AIProviderTypeOpenai {
|
||||
continue
|
||||
}
|
||||
if provider.ID == uuid.Nil || candidate.CreatedAt.After(provider.CreatedAt) {
|
||||
@@ -176,13 +177,12 @@ func ChatModelConfig(t testing.TB, db database.Store, seed database.ChatModelCon
|
||||
}
|
||||
if provider.ID == uuid.Nil {
|
||||
provider = AIProvider(t, db, database.AIProvider{
|
||||
Type: database.AIProviderType(providerName),
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
})
|
||||
}
|
||||
aiProviderID = uuid.NullUUID{UUID: provider.ID, Valid: true}
|
||||
}
|
||||
params := database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
Model: takeFirst(seed.Model, "gpt-4o-mini"),
|
||||
DisplayName: takeFirst(seed.DisplayName, "Test Model"),
|
||||
CreatedBy: seed.CreatedBy,
|
||||
|
||||
@@ -296,7 +296,9 @@ func TestGenerator(t *testing.T) {
|
||||
// Defaults.
|
||||
cfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{})
|
||||
require.NotEqual(t, uuid.Nil, cfg.ID)
|
||||
require.Equal(t, "openai", cfg.Provider)
|
||||
prov, err := db.GetAIProviderByID(context.Background(), cfg.AIProviderID.UUID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "openai", string(prov.Type))
|
||||
require.Equal(t, "gpt-4o-mini", cfg.Model)
|
||||
require.Equal(t, "Test Model", cfg.DisplayName)
|
||||
require.True(t, cfg.Enabled)
|
||||
@@ -304,13 +306,15 @@ func TestGenerator(t *testing.T) {
|
||||
require.Equal(t, int32(70), cfg.CompressionThreshold)
|
||||
|
||||
// Overrides.
|
||||
_ = dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "anthropic"})
|
||||
anthropicProvider := dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "anthropic"})
|
||||
cfg2 := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: uuid.NullUUID{UUID: anthropicProvider.ID, Valid: true},
|
||||
Model: "claude-4",
|
||||
ContextLimit: 200000,
|
||||
})
|
||||
require.Equal(t, "anthropic", cfg2.Provider)
|
||||
prov2, err := db.GetAIProviderByID(context.Background(), cfg2.AIProviderID.UUID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "anthropic", string(prov2.Type))
|
||||
require.Equal(t, "claude-4", cfg2.Model)
|
||||
require.Equal(t, int64(200000), cfg2.ContextLimit)
|
||||
})
|
||||
@@ -325,7 +329,7 @@ func TestGenerator(t *testing.T) {
|
||||
OrganizationID: o.ID,
|
||||
})
|
||||
p := dbgen.ChatProvider(t, db, database.ChatProvider{})
|
||||
m := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{Provider: p.Provider})
|
||||
m := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{AIProviderID: uuid.NullUUID{UUID: p.ID, Valid: true}})
|
||||
|
||||
// Defaults.
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
@@ -360,7 +364,7 @@ func TestGenerator(t *testing.T) {
|
||||
OrganizationID: o.ID,
|
||||
})
|
||||
p := dbgen.ChatProvider(t, db, database.ChatProvider{})
|
||||
m := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{Provider: p.Provider})
|
||||
m := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{AIProviderID: uuid.NullUUID{UUID: p.ID, Valid: true}})
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OwnerID: u.ID,
|
||||
OrganizationID: o.ID,
|
||||
|
||||
+1
-18
@@ -5,7 +5,6 @@ package dbmetrics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"slices"
|
||||
"time"
|
||||
@@ -186,14 +185,6 @@ func (m queryMetricsStore) AutoArchiveInactiveChats(ctx context.Context, arg dat
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) BackfillChatModelConfigProvider(ctx context.Context, arg database.BackfillChatModelConfigProviderParams) (sql.Result, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.BackfillChatModelConfigProvider(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("BackfillChatModelConfigProvider").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "BackfillChatModelConfigProvider").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) BackoffChatDiffStatus(ctx context.Context, arg database.BackoffChatDiffStatusParams) error {
|
||||
start := time.Now()
|
||||
r0 := m.s.BackoffChatDiffStatus(ctx, arg)
|
||||
@@ -538,14 +529,6 @@ func (m queryMetricsStore) DeleteChatModelConfigsByAIProviderID(ctx context.Cont
|
||||
return r0
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) DeleteChatModelConfigsByProvider(ctx context.Context, provider string) error {
|
||||
start := time.Now()
|
||||
r0 := m.s.DeleteChatModelConfigsByProvider(ctx, provider)
|
||||
m.queryLatencies.WithLabelValues("DeleteChatModelConfigsByProvider").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "DeleteChatModelConfigsByProvider").Inc()
|
||||
return r0
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) DeleteChatQueuedMessage(ctx context.Context, arg database.DeleteChatQueuedMessageParams) error {
|
||||
start := time.Now()
|
||||
r0 := m.s.DeleteChatQueuedMessage(ctx, arg)
|
||||
@@ -2018,7 +2001,7 @@ func (m queryMetricsStore) GetEnabledChatModelConfigByID(ctx context.Context, id
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetEnabledChatModelConfigs(ctx context.Context) ([]database.ChatModelConfig, error) {
|
||||
func (m queryMetricsStore) GetEnabledChatModelConfigs(ctx context.Context) ([]database.GetEnabledChatModelConfigsRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetEnabledChatModelConfigs(ctx)
|
||||
m.queryLatencies.WithLabelValues("GetEnabledChatModelConfigs").Observe(time.Since(start).Seconds())
|
||||
|
||||
Generated
+2
-32
@@ -11,7 +11,6 @@ package dbmock
|
||||
|
||||
import (
|
||||
context "context"
|
||||
sql "database/sql"
|
||||
json "encoding/json"
|
||||
reflect "reflect"
|
||||
time "time"
|
||||
@@ -194,21 +193,6 @@ func (mr *MockStoreMockRecorder) AutoArchiveInactiveChats(ctx, arg any) *gomock.
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "AutoArchiveInactiveChats", reflect.TypeOf((*MockStore)(nil).AutoArchiveInactiveChats), ctx, arg)
|
||||
}
|
||||
|
||||
// BackfillChatModelConfigProvider mocks base method.
|
||||
func (m *MockStore) BackfillChatModelConfigProvider(ctx context.Context, arg database.BackfillChatModelConfigProviderParams) (sql.Result, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "BackfillChatModelConfigProvider", ctx, arg)
|
||||
ret0, _ := ret[0].(sql.Result)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// BackfillChatModelConfigProvider indicates an expected call of BackfillChatModelConfigProvider.
|
||||
func (mr *MockStoreMockRecorder) BackfillChatModelConfigProvider(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BackfillChatModelConfigProvider", reflect.TypeOf((*MockStore)(nil).BackfillChatModelConfigProvider), ctx, arg)
|
||||
}
|
||||
|
||||
// BackoffChatDiffStatus mocks base method.
|
||||
func (m *MockStore) BackoffChatDiffStatus(ctx context.Context, arg database.BackoffChatDiffStatusParams) error {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -877,20 +861,6 @@ func (mr *MockStoreMockRecorder) DeleteChatModelConfigsByAIProviderID(ctx, aiPro
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteChatModelConfigsByAIProviderID", reflect.TypeOf((*MockStore)(nil).DeleteChatModelConfigsByAIProviderID), ctx, aiProviderID)
|
||||
}
|
||||
|
||||
// DeleteChatModelConfigsByProvider mocks base method.
|
||||
func (m *MockStore) DeleteChatModelConfigsByProvider(ctx context.Context, provider string) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "DeleteChatModelConfigsByProvider", ctx, provider)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// DeleteChatModelConfigsByProvider indicates an expected call of DeleteChatModelConfigsByProvider.
|
||||
func (mr *MockStoreMockRecorder) DeleteChatModelConfigsByProvider(ctx, provider any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteChatModelConfigsByProvider", reflect.TypeOf((*MockStore)(nil).DeleteChatModelConfigsByProvider), ctx, provider)
|
||||
}
|
||||
|
||||
// DeleteChatQueuedMessage mocks base method.
|
||||
func (m *MockStore) DeleteChatQueuedMessage(ctx context.Context, arg database.DeleteChatQueuedMessageParams) error {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -3733,10 +3703,10 @@ func (mr *MockStoreMockRecorder) GetEnabledChatModelConfigByID(ctx, id any) *gom
|
||||
}
|
||||
|
||||
// GetEnabledChatModelConfigs mocks base method.
|
||||
func (m *MockStore) GetEnabledChatModelConfigs(ctx context.Context) ([]database.ChatModelConfig, error) {
|
||||
func (m *MockStore) GetEnabledChatModelConfigs(ctx context.Context) ([]database.GetEnabledChatModelConfigsRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetEnabledChatModelConfigs", ctx)
|
||||
ret0, _ := ret[0].([]database.ChatModelConfig)
|
||||
ret0, _ := ret[0].([]database.GetEnabledChatModelConfigsRow)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
@@ -2084,7 +2084,6 @@ func TestPurgeChatDebugRuns(t *testing.T) {
|
||||
DisplayName: "OpenAI",
|
||||
})
|
||||
modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
ContextLimit: 8192,
|
||||
})
|
||||
@@ -2311,7 +2310,6 @@ func TestDeleteOldChatFiles(t *testing.T) {
|
||||
DisplayName: "OpenAI",
|
||||
})
|
||||
mc := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
ContextLimit: 8192,
|
||||
})
|
||||
|
||||
Generated
-5
@@ -1909,7 +1909,6 @@ ALTER SEQUENCE chat_messages_id_seq OWNED BY chat_messages.id;
|
||||
|
||||
CREATE TABLE chat_model_configs (
|
||||
id uuid DEFAULT gen_random_uuid() NOT NULL,
|
||||
provider text NOT NULL,
|
||||
model text NOT NULL,
|
||||
display_name text DEFAULT ''::text NOT NULL,
|
||||
created_by uuid,
|
||||
@@ -4625,10 +4624,6 @@ CREATE INDEX idx_chat_model_configs_ai_provider_id ON chat_model_configs USING b
|
||||
|
||||
CREATE INDEX idx_chat_model_configs_enabled ON chat_model_configs USING btree (enabled);
|
||||
|
||||
CREATE INDEX idx_chat_model_configs_provider ON chat_model_configs USING btree (provider);
|
||||
|
||||
CREATE INDEX idx_chat_model_configs_provider_model ON chat_model_configs USING btree (provider, model);
|
||||
|
||||
CREATE UNIQUE INDEX idx_chat_model_configs_single_default ON chat_model_configs USING btree ((1)) WHERE ((is_default = true) AND (deleted = false));
|
||||
|
||||
CREATE INDEX idx_chat_queued_messages_chat_id ON chat_queued_messages USING btree (chat_id);
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
ALTER TABLE chat_model_configs ADD COLUMN provider text;
|
||||
|
||||
UPDATE chat_model_configs cmc
|
||||
SET provider = ap.type::text
|
||||
FROM ai_providers ap
|
||||
WHERE ap.id = cmc.ai_provider_id;
|
||||
|
||||
UPDATE chat_model_configs SET provider = '' WHERE provider IS NULL;
|
||||
|
||||
ALTER TABLE chat_model_configs ALTER COLUMN provider SET NOT NULL;
|
||||
|
||||
CREATE INDEX idx_chat_model_configs_provider ON chat_model_configs USING btree (provider);
|
||||
CREATE INDEX idx_chat_model_configs_provider_model ON chat_model_configs USING btree (provider, model);
|
||||
@@ -0,0 +1,4 @@
|
||||
DROP INDEX idx_chat_model_configs_provider;
|
||||
DROP INDEX idx_chat_model_configs_provider_model;
|
||||
|
||||
ALTER TABLE chat_model_configs DROP COLUMN provider;
|
||||
+41
@@ -0,0 +1,41 @@
|
||||
-- Soft-deleted chat model config whose provider never had an ai_providers
|
||||
-- backfill match, so it reaches later migrations with ai_provider_id IS NULL.
|
||||
--
|
||||
-- This row exercises the 000534 down migration's
|
||||
-- `UPDATE ... SET provider = '' WHERE provider IS NULL` sweep: its NULL
|
||||
-- ai_provider_id means the backfill join leaves provider NULL, and the sweep
|
||||
-- must populate it before `ALTER COLUMN provider SET NOT NULL`.
|
||||
--
|
||||
-- It is inserted at version 000475 (after 000474 dropped the provider foreign
|
||||
-- key) so the provider value need not reference a chat_providers row, and the
|
||||
-- 000504/000505 backfill (which matches on `cmc.provider = cp.provider`) skips
|
||||
-- it. `deleted = TRUE` keeps it out of idx_chat_model_configs_single_default
|
||||
-- and satisfies chat_model_configs_ai_provider_required_when_active (added in
|
||||
-- 000505), which permits a NULL ai_provider_id only for deleted rows.
|
||||
INSERT INTO chat_model_configs (
|
||||
id,
|
||||
provider,
|
||||
model,
|
||||
display_name,
|
||||
enabled,
|
||||
is_default,
|
||||
deleted,
|
||||
deleted_at,
|
||||
context_limit,
|
||||
compression_threshold,
|
||||
created_at,
|
||||
updated_at
|
||||
) VALUES (
|
||||
'b3a1d2c4-5e6f-4a7b-8c9d-0e1f2a3b4c5d',
|
||||
'legacy-removed',
|
||||
'legacy-model',
|
||||
'Legacy Soft Deleted',
|
||||
FALSE,
|
||||
FALSE,
|
||||
TRUE,
|
||||
'2024-01-01 00:00:00+00',
|
||||
200000,
|
||||
70,
|
||||
'2024-01-01 00:00:00+00',
|
||||
'2024-01-01 00:00:00+00'
|
||||
);
|
||||
Generated
-1
@@ -4953,7 +4953,6 @@ type ChatMessage struct {
|
||||
|
||||
type ChatModelConfig struct {
|
||||
ID uuid.UUID `db:"id" json:"id"`
|
||||
Provider string `db:"provider" json:"provider"`
|
||||
Model string `db:"model" json:"model"`
|
||||
DisplayName string `db:"display_name" json:"display_name"`
|
||||
CreatedBy uuid.NullUUID `db:"created_by" json:"created_by"`
|
||||
|
||||
Generated
+2
-8
@@ -6,7 +6,6 @@ package database
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"time"
|
||||
|
||||
@@ -71,11 +70,6 @@ type sqlcQuerier interface {
|
||||
// created_at ASC flows through to dbpurge's digest truncation; see
|
||||
// buildDigestData in dbpurge.go for the tradeoff rationale.
|
||||
AutoArchiveInactiveChats(ctx context.Context, arg AutoArchiveInactiveChatsParams) ([]AutoArchiveInactiveChatsRow, error)
|
||||
// old_provider is matched as text; new_provider is also cast to ai_provider_type
|
||||
// for the EXISTS check against ai_providers.type.
|
||||
// ai_provider_id IS NOT NULL is defensive; the check constraint already
|
||||
// enforces that non-deleted rows always have a provider ID.
|
||||
BackfillChatModelConfigProvider(ctx context.Context, arg BackfillChatModelConfigProviderParams) (sql.Result, error)
|
||||
BackoffChatDiffStatus(ctx context.Context, arg BackoffChatDiffStatusParams) error
|
||||
// Deletes heartbeat rows for the supplied (chat_id, runner_id) pairs.
|
||||
BatchDeleteChatHeartbeats(ctx context.Context, arg BatchDeleteChatHeartbeatsParams) (int64, error)
|
||||
@@ -150,7 +144,6 @@ type sqlcQuerier interface {
|
||||
DeleteChatDebugDataByChatID(ctx context.Context, arg DeleteChatDebugDataByChatIDParams) (int64, error)
|
||||
DeleteChatModelConfigByID(ctx context.Context, id uuid.UUID) error
|
||||
DeleteChatModelConfigsByAIProviderID(ctx context.Context, aiProviderID uuid.UUID) error
|
||||
DeleteChatModelConfigsByProvider(ctx context.Context, provider string) error
|
||||
DeleteChatQueuedMessage(ctx context.Context, arg DeleteChatQueuedMessageParams) error
|
||||
// Deletes a queued message, scoped to the parent chat. Returns the
|
||||
// number of affected rows so callers can detect missing rows without
|
||||
@@ -449,6 +442,7 @@ type sqlcQuerier interface {
|
||||
GetChatModelConfigByID(ctx context.Context, id uuid.UUID) (ChatModelConfig, error)
|
||||
GetChatModelConfigs(ctx context.Context) ([]ChatModelConfig, error)
|
||||
// Returns all model configurations for telemetry snapshot collection.
|
||||
// deleted = false guarantees ai_provider_id is non-null, so INNER JOIN is safe.
|
||||
GetChatModelConfigsForTelemetry(ctx context.Context) ([]GetChatModelConfigsForTelemetryRow, error)
|
||||
// GetChatPersonalModelOverridesEnabled returns whether users may configure
|
||||
// personal chat model overrides. It defaults to false when unset.
|
||||
@@ -539,7 +533,7 @@ type sqlcQuerier interface {
|
||||
// Providers can be disabled independently of their model configs.
|
||||
// Check both to ensure the selected config is actually usable.
|
||||
GetEnabledChatModelConfigByID(ctx context.Context, id uuid.UUID) (ChatModelConfig, error)
|
||||
GetEnabledChatModelConfigs(ctx context.Context) ([]ChatModelConfig, error)
|
||||
GetEnabledChatModelConfigs(ctx context.Context) ([]GetEnabledChatModelConfigsRow, error)
|
||||
GetEnabledMCPServerConfigs(ctx context.Context) ([]MCPServerConfig, error)
|
||||
// GetExternalAgentTokensByTemplateID returns the auth tokens for all
|
||||
// non-deleted external agents on the latest build of every running workspace
|
||||
|
||||
@@ -1250,7 +1250,6 @@ func TestChatContextHydration(t *testing.T) {
|
||||
owner := dbgen.User(t, db, database.User{})
|
||||
_ = dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "openai", DisplayName: "OpenAI"})
|
||||
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -1381,7 +1380,6 @@ func TestGetAuthorizedChats(t *testing.T) {
|
||||
DisplayName: "OpenAI",
|
||||
})
|
||||
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -1645,7 +1643,6 @@ func TestGetAuthorizedChatsACLSharing(t *testing.T) {
|
||||
|
||||
dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "openai", DisplayName: "OpenAI"})
|
||||
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -1766,7 +1763,6 @@ func TestGetAuthorizedChatsACLSharingGroupACL(t *testing.T) {
|
||||
|
||||
dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "openai", DisplayName: "OpenAI"})
|
||||
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -1869,7 +1865,6 @@ func TestGetAuthorizedChatsByChatFileIDACLSharing(t *testing.T) {
|
||||
|
||||
dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "openai", DisplayName: "OpenAI"})
|
||||
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -11097,24 +11092,21 @@ func TestGetEnabledChatModelConfigsUsesAIProviders(t *testing.T) {
|
||||
params.Enabled = false
|
||||
})
|
||||
enabledConfig := dbgen.ChatModelConfig(t, store, database.ChatModelConfig{
|
||||
Provider: string(enabledProvider.Type),
|
||||
Model: "openrouter-model-" + uuid.NewString(),
|
||||
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(),
|
||||
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(),
|
||||
Model: "disabled-model-" + uuid.NewString(),
|
||||
AIProviderID: uuid.NullUUID{
|
||||
UUID: enabledProvider.ID,
|
||||
Valid: true,
|
||||
@@ -11125,14 +11117,14 @@ func TestGetEnabledChatModelConfigsUsesAIProviders(t *testing.T) {
|
||||
|
||||
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(row database.GetEnabledChatModelConfigsRow) bool {
|
||||
return row.ChatModelConfig.ID == enabledConfig.ID
|
||||
}))
|
||||
require.False(t, slices.ContainsFunc(configs, func(config database.ChatModelConfig) bool {
|
||||
return config.ID == disabledProviderConfig.ID
|
||||
require.False(t, slices.ContainsFunc(configs, func(row database.GetEnabledChatModelConfigsRow) bool {
|
||||
return row.ChatModelConfig.ID == disabledProviderConfig.ID
|
||||
}))
|
||||
require.False(t, slices.ContainsFunc(configs, func(config database.ChatModelConfig) bool {
|
||||
return config.ID == disabledModelConfig.ID
|
||||
require.False(t, slices.ContainsFunc(configs, func(row database.GetEnabledChatModelConfigsRow) bool {
|
||||
return row.ChatModelConfig.ID == disabledModelConfig.ID
|
||||
}))
|
||||
|
||||
config, err := store.GetEnabledChatModelConfigByID(ctx, enabledConfig.ID)
|
||||
@@ -11150,16 +11142,16 @@ func insertChatModelConfigForTest(
|
||||
ctx context.Context,
|
||||
t testing.TB,
|
||||
store database.Store,
|
||||
providerType string,
|
||||
params database.InsertChatModelConfigParams,
|
||||
) (database.ChatModelConfig, error) {
|
||||
t.Helper()
|
||||
if params.AIProviderID.Valid {
|
||||
return store.InsertChatModelConfig(ctx, params)
|
||||
}
|
||||
providerName := params.Provider
|
||||
providerName := providerType
|
||||
if providerName == "" {
|
||||
providerName = "openai"
|
||||
params.Provider = providerName
|
||||
}
|
||||
providers, err := store.GetAIProviders(ctx, database.GetAIProvidersParams{IncludeDisabled: true})
|
||||
if err != nil {
|
||||
@@ -11198,8 +11190,7 @@ func TestInsertChatMessages(t *testing.T) {
|
||||
) database.ChatModelConfig {
|
||||
t.Helper()
|
||||
|
||||
modelConfig, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: provider,
|
||||
modelConfig, err := insertChatModelConfigForTest(ctx, t, store, provider, database.InsertChatModelConfigParams{
|
||||
Model: model,
|
||||
DisplayName: displayName,
|
||||
CreatedBy: uuid.NullUUID{UUID: userID, Valid: true},
|
||||
@@ -11410,8 +11401,7 @@ func TestGetChatMessagesForPromptByChatID(t *testing.T) {
|
||||
APIKey: "test-key",
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, db, database.InsertChatModelConfigParams{
|
||||
Provider: "openai",
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, db, "openai", database.InsertChatModelConfigParams{
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
Model: "test-model",
|
||||
DisplayName: "Test Model",
|
||||
@@ -11786,8 +11776,7 @@ func TestChatPinOrderQueries(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(bg, t, db, database.InsertChatModelConfigParams{
|
||||
Provider: "openai",
|
||||
modelCfg, err := insertChatModelConfigForTest(bg, t, db, "openai", database.InsertChatModelConfigParams{
|
||||
Model: "test-model",
|
||||
DisplayName: "Test Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -11966,8 +11955,7 @@ func TestChatPinOrderConstraints(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(bg, t, db, database.InsertChatModelConfigParams{
|
||||
Provider: "openai",
|
||||
modelCfg, err := insertChatModelConfigForTest(bg, t, db, "openai", database.InsertChatModelConfigParams{
|
||||
Model: "test-model",
|
||||
DisplayName: "Test Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -12058,8 +12046,7 @@ func TestChatLabels(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, db, database.InsertChatModelConfigParams{
|
||||
Provider: "openai",
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, db, "openai", database.InsertChatModelConfigParams{
|
||||
Model: "test-model",
|
||||
DisplayName: "Test Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -12357,8 +12344,7 @@ func TestUpdateChatLastTurnSummary(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, db, database.InsertChatModelConfigParams{
|
||||
Provider: "openai",
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, db, "openai", database.InsertChatModelConfigParams{
|
||||
Model: "test-model",
|
||||
DisplayName: "Test Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -12494,8 +12480,7 @@ func TestDeleteChatDebugDataAfterMessageIDIncludesTriggeredRuns(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
@@ -12686,8 +12671,7 @@ func TestDeleteChatDebugDataAfterMessageIDStepLevelFieldBoundariesAndNulls(t *te
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
@@ -12943,8 +12927,7 @@ func TestFinalizeStaleChatDebugRows(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
@@ -13381,8 +13364,7 @@ func TestChatDebugSQLGuards(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
@@ -13514,8 +13496,7 @@ func TestChatDebugRunCOALESCEPreservation(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
@@ -13628,8 +13609,7 @@ func TestChatDebugStepCOALESCEPreservation(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
@@ -13752,8 +13732,7 @@ func TestDeleteChatDebugDataAfterMessageIDNullMessagesSurvive(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
@@ -13849,8 +13828,7 @@ func TestDeleteChatDebugDataAfterMessageIDStartedBeforeFiltersNewerRuns(t *testi
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
@@ -13960,8 +13938,7 @@ func TestDeleteChatDebugDataByChatIDStartedBeforeFiltersNewerRuns(t *testing.T)
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: providerName,
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{
|
||||
Model: modelName,
|
||||
DisplayName: "Debug Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
@@ -14045,7 +14022,6 @@ func TestGetChatsFilter(t *testing.T) {
|
||||
}, "test-key")
|
||||
|
||||
modelCfg, err := store.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
|
||||
Provider: "openai",
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
Model: "test-model-" + uuid.NewString(),
|
||||
DisplayName: "Test Model",
|
||||
@@ -14343,8 +14319,7 @@ func TestChatHasUnread(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{
|
||||
Provider: "openai",
|
||||
modelCfg, err := insertChatModelConfigForTest(ctx, t, store, "openai", database.InsertChatModelConfigParams{
|
||||
Model: "test-model-" + uuid.NewString(),
|
||||
DisplayName: "Test Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
|
||||
Generated
+65
-114
@@ -5025,37 +5025,6 @@ func (q *sqlQuerier) InsertChatFile(ctx context.Context, arg InsertChatFileParam
|
||||
return i, err
|
||||
}
|
||||
|
||||
const backfillChatModelConfigProvider = `-- name: BackfillChatModelConfigProvider :execresult
|
||||
UPDATE
|
||||
chat_model_configs
|
||||
SET
|
||||
provider = $1::text,
|
||||
updated_at = NOW()
|
||||
WHERE
|
||||
provider = $2::text
|
||||
AND deleted = FALSE
|
||||
AND ai_provider_id IS NOT NULL
|
||||
AND EXISTS (
|
||||
SELECT 1 FROM ai_providers
|
||||
WHERE id = chat_model_configs.ai_provider_id
|
||||
AND type = $1::ai_provider_type
|
||||
AND deleted = FALSE
|
||||
)
|
||||
`
|
||||
|
||||
type BackfillChatModelConfigProviderParams struct {
|
||||
NewProvider string `db:"new_provider" json:"new_provider"`
|
||||
OldProvider string `db:"old_provider" json:"old_provider"`
|
||||
}
|
||||
|
||||
// old_provider is matched as text; new_provider is also cast to ai_provider_type
|
||||
// for the EXISTS check against ai_providers.type.
|
||||
// ai_provider_id IS NOT NULL is defensive; the check constraint already
|
||||
// enforces that non-deleted rows always have a provider ID.
|
||||
func (q *sqlQuerier) BackfillChatModelConfigProvider(ctx context.Context, arg BackfillChatModelConfigProviderParams) (sql.Result, error) {
|
||||
return q.db.ExecContext(ctx, backfillChatModelConfigProvider, arg.NewProvider, arg.OldProvider)
|
||||
}
|
||||
|
||||
const deleteChatModelConfigByID = `-- name: DeleteChatModelConfigByID :exec
|
||||
UPDATE
|
||||
chat_model_configs
|
||||
@@ -5089,26 +5058,9 @@ func (q *sqlQuerier) DeleteChatModelConfigsByAIProviderID(ctx context.Context, a
|
||||
return err
|
||||
}
|
||||
|
||||
const deleteChatModelConfigsByProvider = `-- name: DeleteChatModelConfigsByProvider :exec
|
||||
UPDATE
|
||||
chat_model_configs
|
||||
SET
|
||||
deleted = TRUE,
|
||||
deleted_at = NOW(),
|
||||
updated_at = NOW()
|
||||
WHERE
|
||||
provider = $1::text
|
||||
AND deleted = FALSE
|
||||
`
|
||||
|
||||
func (q *sqlQuerier) DeleteChatModelConfigsByProvider(ctx context.Context, provider string) error {
|
||||
_, err := q.db.ExecContext(ctx, deleteChatModelConfigsByProvider, provider)
|
||||
return err
|
||||
}
|
||||
|
||||
const getChatModelConfigByID = `-- name: GetChatModelConfigByID :one
|
||||
SELECT
|
||||
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
|
||||
id, 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
|
||||
FROM
|
||||
chat_model_configs
|
||||
WHERE
|
||||
@@ -5121,7 +5073,6 @@ func (q *sqlQuerier) GetChatModelConfigByID(ctx context.Context, id uuid.UUID) (
|
||||
var i ChatModelConfig
|
||||
err := row.Scan(
|
||||
&i.ID,
|
||||
&i.Provider,
|
||||
&i.Model,
|
||||
&i.DisplayName,
|
||||
&i.CreatedBy,
|
||||
@@ -5142,16 +5093,18 @@ func (q *sqlQuerier) GetChatModelConfigByID(ctx context.Context, id uuid.UUID) (
|
||||
|
||||
const getChatModelConfigs = `-- name: GetChatModelConfigs :many
|
||||
SELECT
|
||||
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
|
||||
cmc.id, 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
|
||||
chat_model_configs cmc
|
||||
LEFT JOIN
|
||||
ai_providers ap ON ap.id = cmc.ai_provider_id
|
||||
WHERE
|
||||
deleted = FALSE
|
||||
cmc.deleted = FALSE
|
||||
ORDER BY
|
||||
provider ASC,
|
||||
model ASC,
|
||||
updated_at DESC,
|
||||
id DESC
|
||||
ap.type::text ASC,
|
||||
cmc.model ASC,
|
||||
cmc.updated_at DESC,
|
||||
cmc.id DESC
|
||||
`
|
||||
|
||||
func (q *sqlQuerier) GetChatModelConfigs(ctx context.Context) ([]ChatModelConfig, error) {
|
||||
@@ -5165,7 +5118,6 @@ func (q *sqlQuerier) GetChatModelConfigs(ctx context.Context) ([]ChatModelConfig
|
||||
var i ChatModelConfig
|
||||
if err := rows.Scan(
|
||||
&i.ID,
|
||||
&i.Provider,
|
||||
&i.Model,
|
||||
&i.DisplayName,
|
||||
&i.CreatedBy,
|
||||
@@ -5196,7 +5148,7 @@ func (q *sqlQuerier) GetChatModelConfigs(ctx context.Context) ([]ChatModelConfig
|
||||
|
||||
const getDefaultChatModelConfig = `-- name: GetDefaultChatModelConfig :one
|
||||
SELECT
|
||||
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
|
||||
id, 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
|
||||
FROM
|
||||
chat_model_configs
|
||||
WHERE
|
||||
@@ -5209,7 +5161,6 @@ func (q *sqlQuerier) GetDefaultChatModelConfig(ctx context.Context) (ChatModelCo
|
||||
var i ChatModelConfig
|
||||
err := row.Scan(
|
||||
&i.ID,
|
||||
&i.Provider,
|
||||
&i.Model,
|
||||
&i.DisplayName,
|
||||
&i.CreatedBy,
|
||||
@@ -5230,7 +5181,7 @@ func (q *sqlQuerier) GetDefaultChatModelConfig(ctx context.Context) (ChatModelCo
|
||||
|
||||
const getEnabledChatModelConfigByID = `-- name: GetEnabledChatModelConfigByID :one
|
||||
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
|
||||
cmc.id, 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
|
||||
@@ -5250,7 +5201,6 @@ func (q *sqlQuerier) GetEnabledChatModelConfigByID(ctx context.Context, id uuid.
|
||||
var i ChatModelConfig
|
||||
err := row.Scan(
|
||||
&i.ID,
|
||||
&i.Provider,
|
||||
&i.Model,
|
||||
&i.DisplayName,
|
||||
&i.CreatedBy,
|
||||
@@ -5271,7 +5221,8 @@ func (q *sqlQuerier) GetEnabledChatModelConfigByID(ctx context.Context, id uuid.
|
||||
|
||||
const getEnabledChatModelConfigs = `-- name: GetEnabledChatModelConfigs :many
|
||||
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
|
||||
cmc.id, 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,
|
||||
ap.type::text AS provider
|
||||
FROM
|
||||
chat_model_configs cmc
|
||||
JOIN
|
||||
@@ -5282,38 +5233,43 @@ WHERE
|
||||
AND ap.enabled = TRUE
|
||||
AND ap.deleted = FALSE
|
||||
ORDER BY
|
||||
cmc.provider ASC,
|
||||
ap.type::text ASC,
|
||||
cmc.model ASC,
|
||||
cmc.updated_at DESC,
|
||||
cmc.id DESC
|
||||
`
|
||||
|
||||
func (q *sqlQuerier) GetEnabledChatModelConfigs(ctx context.Context) ([]ChatModelConfig, error) {
|
||||
type GetEnabledChatModelConfigsRow struct {
|
||||
ChatModelConfig ChatModelConfig `db:"chat_model_config" json:"chat_model_config"`
|
||||
Provider string `db:"provider" json:"provider"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) GetEnabledChatModelConfigs(ctx context.Context) ([]GetEnabledChatModelConfigsRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getEnabledChatModelConfigs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var items []ChatModelConfig
|
||||
var items []GetEnabledChatModelConfigsRow
|
||||
for rows.Next() {
|
||||
var i ChatModelConfig
|
||||
var i GetEnabledChatModelConfigsRow
|
||||
if err := rows.Scan(
|
||||
&i.ID,
|
||||
&i.ChatModelConfig.ID,
|
||||
&i.ChatModelConfig.Model,
|
||||
&i.ChatModelConfig.DisplayName,
|
||||
&i.ChatModelConfig.CreatedBy,
|
||||
&i.ChatModelConfig.UpdatedBy,
|
||||
&i.ChatModelConfig.Enabled,
|
||||
&i.ChatModelConfig.IsDefault,
|
||||
&i.ChatModelConfig.Deleted,
|
||||
&i.ChatModelConfig.DeletedAt,
|
||||
&i.ChatModelConfig.CreatedAt,
|
||||
&i.ChatModelConfig.UpdatedAt,
|
||||
&i.ChatModelConfig.ContextLimit,
|
||||
&i.ChatModelConfig.CompressionThreshold,
|
||||
&i.ChatModelConfig.Options,
|
||||
&i.ChatModelConfig.AIProviderID,
|
||||
&i.Provider,
|
||||
&i.Model,
|
||||
&i.DisplayName,
|
||||
&i.CreatedBy,
|
||||
&i.UpdatedBy,
|
||||
&i.Enabled,
|
||||
&i.IsDefault,
|
||||
&i.Deleted,
|
||||
&i.DeletedAt,
|
||||
&i.CreatedAt,
|
||||
&i.UpdatedAt,
|
||||
&i.ContextLimit,
|
||||
&i.CompressionThreshold,
|
||||
&i.Options,
|
||||
&i.AIProviderID,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -5330,7 +5286,6 @@ func (q *sqlQuerier) GetEnabledChatModelConfigs(ctx context.Context) ([]ChatMode
|
||||
|
||||
const insertChatModelConfig = `-- name: InsertChatModelConfig :one
|
||||
INSERT INTO chat_model_configs (
|
||||
provider,
|
||||
model,
|
||||
display_name,
|
||||
created_by,
|
||||
@@ -5344,22 +5299,20 @@ INSERT INTO chat_model_configs (
|
||||
) VALUES (
|
||||
$1::text,
|
||||
$2::text,
|
||||
$3::text,
|
||||
$3::uuid,
|
||||
$4::uuid,
|
||||
$5::uuid,
|
||||
$5::boolean,
|
||||
$6::boolean,
|
||||
$7::boolean,
|
||||
$8::bigint,
|
||||
$9::integer,
|
||||
$10::jsonb,
|
||||
$11::uuid
|
||||
$7::bigint,
|
||||
$8::integer,
|
||||
$9::jsonb,
|
||||
$10::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
|
||||
id, 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
|
||||
`
|
||||
|
||||
type InsertChatModelConfigParams struct {
|
||||
Provider string `db:"provider" json:"provider"`
|
||||
Model string `db:"model" json:"model"`
|
||||
DisplayName string `db:"display_name" json:"display_name"`
|
||||
CreatedBy uuid.NullUUID `db:"created_by" json:"created_by"`
|
||||
@@ -5374,7 +5327,6 @@ type InsertChatModelConfigParams struct {
|
||||
|
||||
func (q *sqlQuerier) InsertChatModelConfig(ctx context.Context, arg InsertChatModelConfigParams) (ChatModelConfig, error) {
|
||||
row := q.db.QueryRowContext(ctx, insertChatModelConfig,
|
||||
arg.Provider,
|
||||
arg.Model,
|
||||
arg.DisplayName,
|
||||
arg.CreatedBy,
|
||||
@@ -5389,7 +5341,6 @@ func (q *sqlQuerier) InsertChatModelConfig(ctx context.Context, arg InsertChatMo
|
||||
var i ChatModelConfig
|
||||
err := row.Scan(
|
||||
&i.ID,
|
||||
&i.Provider,
|
||||
&i.Model,
|
||||
&i.DisplayName,
|
||||
&i.CreatedBy,
|
||||
@@ -5428,26 +5379,24 @@ const updateChatModelConfig = `-- name: UpdateChatModelConfig :one
|
||||
UPDATE
|
||||
chat_model_configs
|
||||
SET
|
||||
provider = $1::text,
|
||||
model = $2::text,
|
||||
display_name = $3::text,
|
||||
updated_by = $4::uuid,
|
||||
enabled = $5::boolean,
|
||||
is_default = $6::boolean,
|
||||
context_limit = $7::bigint,
|
||||
compression_threshold = $8::integer,
|
||||
options = $9::jsonb,
|
||||
ai_provider_id = $10::uuid,
|
||||
model = $1::text,
|
||||
display_name = $2::text,
|
||||
updated_by = $3::uuid,
|
||||
enabled = $4::boolean,
|
||||
is_default = $5::boolean,
|
||||
context_limit = $6::bigint,
|
||||
compression_threshold = $7::integer,
|
||||
options = $8::jsonb,
|
||||
ai_provider_id = $9::uuid,
|
||||
updated_at = NOW()
|
||||
WHERE
|
||||
id = $11::uuid
|
||||
id = $10::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
|
||||
id, 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
|
||||
`
|
||||
|
||||
type UpdateChatModelConfigParams struct {
|
||||
Provider string `db:"provider" json:"provider"`
|
||||
Model string `db:"model" json:"model"`
|
||||
DisplayName string `db:"display_name" json:"display_name"`
|
||||
UpdatedBy uuid.NullUUID `db:"updated_by" json:"updated_by"`
|
||||
@@ -5462,7 +5411,6 @@ type UpdateChatModelConfigParams struct {
|
||||
|
||||
func (q *sqlQuerier) UpdateChatModelConfig(ctx context.Context, arg UpdateChatModelConfigParams) (ChatModelConfig, error) {
|
||||
row := q.db.QueryRowContext(ctx, updateChatModelConfig,
|
||||
arg.Provider,
|
||||
arg.Model,
|
||||
arg.DisplayName,
|
||||
arg.UpdatedBy,
|
||||
@@ -5477,7 +5425,6 @@ func (q *sqlQuerier) UpdateChatModelConfig(ctx context.Context, arg UpdateChatMo
|
||||
var i ChatModelConfig
|
||||
err := row.Scan(
|
||||
&i.ID,
|
||||
&i.Provider,
|
||||
&i.Model,
|
||||
&i.DisplayName,
|
||||
&i.CreatedBy,
|
||||
@@ -6967,7 +6914,7 @@ const getChatCostPerModel = `-- name: GetChatCostPerModel :many
|
||||
SELECT
|
||||
cmc.id AS model_config_id,
|
||||
cmc.display_name,
|
||||
cmc.provider,
|
||||
COALESCE(ap.type::text, '')::text AS provider,
|
||||
cmc.model,
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
@@ -6988,13 +6935,15 @@ JOIN
|
||||
chats c ON c.id = cm.chat_id
|
||||
JOIN
|
||||
chat_model_configs cmc ON cmc.id = cm.model_config_id
|
||||
LEFT JOIN
|
||||
ai_providers ap ON ap.id = cmc.ai_provider_id
|
||||
WHERE
|
||||
c.owner_id = $1::uuid
|
||||
AND cm.role = 'assistant'
|
||||
AND cm.created_at >= $2::timestamptz
|
||||
AND cm.created_at < $3::timestamptz
|
||||
GROUP BY
|
||||
cmc.id, cmc.display_name, cmc.provider, cmc.model
|
||||
cmc.id, cmc.display_name, ap.type, cmc.model
|
||||
ORDER BY
|
||||
total_cost_micros DESC
|
||||
`
|
||||
@@ -7960,9 +7909,10 @@ func (q *sqlQuerier) GetChatMessagesForPromptByChatID(ctx context.Context, chatI
|
||||
}
|
||||
|
||||
const getChatModelConfigsForTelemetry = `-- name: GetChatModelConfigsForTelemetry :many
|
||||
SELECT id, provider, model, context_limit, enabled, is_default
|
||||
FROM chat_model_configs
|
||||
WHERE deleted = false
|
||||
SELECT cmc.id, ap.type::text AS provider, cmc.model, cmc.context_limit, cmc.enabled, cmc.is_default
|
||||
FROM chat_model_configs cmc
|
||||
JOIN ai_providers ap ON ap.id = cmc.ai_provider_id
|
||||
WHERE cmc.deleted = false
|
||||
`
|
||||
|
||||
type GetChatModelConfigsForTelemetryRow struct {
|
||||
@@ -7975,6 +7925,7 @@ type GetChatModelConfigsForTelemetryRow struct {
|
||||
}
|
||||
|
||||
// Returns all model configurations for telemetry snapshot collection.
|
||||
// deleted = false guarantees ai_provider_id is non-null, so INNER JOIN is safe.
|
||||
func (q *sqlQuerier) GetChatModelConfigsForTelemetry(ctx context.Context) ([]GetChatModelConfigsForTelemetryRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, getChatModelConfigsForTelemetry)
|
||||
if err != nil {
|
||||
|
||||
@@ -18,20 +18,23 @@ WHERE
|
||||
|
||||
-- name: GetChatModelConfigs :many
|
||||
SELECT
|
||||
*
|
||||
cmc.*
|
||||
FROM
|
||||
chat_model_configs
|
||||
chat_model_configs cmc
|
||||
LEFT JOIN
|
||||
ai_providers ap ON ap.id = cmc.ai_provider_id
|
||||
WHERE
|
||||
deleted = FALSE
|
||||
cmc.deleted = FALSE
|
||||
ORDER BY
|
||||
provider ASC,
|
||||
model ASC,
|
||||
updated_at DESC,
|
||||
id DESC;
|
||||
ap.type::text ASC,
|
||||
cmc.model ASC,
|
||||
cmc.updated_at DESC,
|
||||
cmc.id DESC;
|
||||
|
||||
-- name: GetEnabledChatModelConfigs :many
|
||||
SELECT
|
||||
cmc.*
|
||||
sqlc.embed(cmc),
|
||||
ap.type::text AS provider
|
||||
FROM
|
||||
chat_model_configs cmc
|
||||
JOIN
|
||||
@@ -42,7 +45,7 @@ WHERE
|
||||
AND ap.enabled = TRUE
|
||||
AND ap.deleted = FALSE
|
||||
ORDER BY
|
||||
cmc.provider ASC,
|
||||
ap.type::text ASC,
|
||||
cmc.model ASC,
|
||||
cmc.updated_at DESC,
|
||||
cmc.id DESC;
|
||||
@@ -65,7 +68,6 @@ WHERE
|
||||
|
||||
-- name: InsertChatModelConfig :one
|
||||
INSERT INTO chat_model_configs (
|
||||
provider,
|
||||
model,
|
||||
display_name,
|
||||
created_by,
|
||||
@@ -77,7 +79,6 @@ INSERT INTO chat_model_configs (
|
||||
options,
|
||||
ai_provider_id
|
||||
) VALUES (
|
||||
@provider::text,
|
||||
@model::text,
|
||||
@display_name::text,
|
||||
sqlc.narg('created_by')::uuid,
|
||||
@@ -96,7 +97,6 @@ RETURNING
|
||||
UPDATE
|
||||
chat_model_configs
|
||||
SET
|
||||
provider = @provider::text,
|
||||
model = @model::text,
|
||||
display_name = @display_name::text,
|
||||
updated_by = sqlc.narg('updated_by')::uuid,
|
||||
@@ -133,38 +133,6 @@ SET
|
||||
WHERE
|
||||
id = @id::uuid;
|
||||
|
||||
-- name: DeleteChatModelConfigsByProvider :exec
|
||||
UPDATE
|
||||
chat_model_configs
|
||||
SET
|
||||
deleted = TRUE,
|
||||
deleted_at = NOW(),
|
||||
updated_at = NOW()
|
||||
WHERE
|
||||
provider = @provider::text
|
||||
AND deleted = FALSE;
|
||||
|
||||
-- name: BackfillChatModelConfigProvider :execresult
|
||||
-- old_provider is matched as text; new_provider is also cast to ai_provider_type
|
||||
-- for the EXISTS check against ai_providers.type.
|
||||
-- ai_provider_id IS NOT NULL is defensive; the check constraint already
|
||||
-- enforces that non-deleted rows always have a provider ID.
|
||||
UPDATE
|
||||
chat_model_configs
|
||||
SET
|
||||
provider = @new_provider::text,
|
||||
updated_at = NOW()
|
||||
WHERE
|
||||
provider = @old_provider::text
|
||||
AND deleted = FALSE
|
||||
AND ai_provider_id IS NOT NULL
|
||||
AND EXISTS (
|
||||
SELECT 1 FROM ai_providers
|
||||
WHERE id = chat_model_configs.ai_provider_id
|
||||
AND type = @new_provider::ai_provider_type
|
||||
AND deleted = FALSE
|
||||
);
|
||||
|
||||
-- name: DeleteChatModelConfigsByAIProviderID :exec
|
||||
UPDATE
|
||||
chat_model_configs
|
||||
|
||||
@@ -2220,7 +2220,7 @@ WHERE
|
||||
SELECT
|
||||
cmc.id AS model_config_id,
|
||||
cmc.display_name,
|
||||
cmc.provider,
|
||||
COALESCE(ap.type::text, '')::text AS provider,
|
||||
cmc.model,
|
||||
COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros,
|
||||
COUNT(*) FILTER (
|
||||
@@ -2241,13 +2241,15 @@ JOIN
|
||||
chats c ON c.id = cm.chat_id
|
||||
JOIN
|
||||
chat_model_configs cmc ON cmc.id = cm.model_config_id
|
||||
LEFT JOIN
|
||||
ai_providers ap ON ap.id = cmc.ai_provider_id
|
||||
WHERE
|
||||
c.owner_id = @owner_id::uuid
|
||||
AND cm.role = 'assistant'
|
||||
AND cm.created_at >= @start_date::timestamptz
|
||||
AND cm.created_at < @end_date::timestamptz
|
||||
GROUP BY
|
||||
cmc.id, cmc.display_name, cmc.provider, cmc.model
|
||||
cmc.id, cmc.display_name, ap.type, cmc.model
|
||||
ORDER BY
|
||||
total_cost_micros DESC;
|
||||
|
||||
@@ -2580,9 +2582,11 @@ GROUP BY cm.chat_id;
|
||||
|
||||
-- name: GetChatModelConfigsForTelemetry :many
|
||||
-- Returns all model configurations for telemetry snapshot collection.
|
||||
SELECT id, provider, model, context_limit, enabled, is_default
|
||||
FROM chat_model_configs
|
||||
WHERE deleted = false;
|
||||
-- deleted = false guarantees ai_provider_id is non-null, so INNER JOIN is safe.
|
||||
SELECT cmc.id, ap.type::text AS provider, cmc.model, cmc.context_limit, cmc.enabled, cmc.is_default
|
||||
FROM chat_model_configs cmc
|
||||
JOIN ai_providers ap ON ap.id = cmc.ai_provider_id
|
||||
WHERE cmc.deleted = false;
|
||||
-- name: GetActiveChatsByAgentID :many
|
||||
SELECT *
|
||||
FROM chats_expanded
|
||||
|
||||
+20
-49
@@ -774,7 +774,7 @@ func (api *API) chatPersonalModelOverrideDeploymentDefaults(
|
||||
type userChatModelAvailability struct {
|
||||
configuredProviders []chatprovider.ConfiguredProvider
|
||||
configuredModels []chatprovider.ConfiguredModel
|
||||
enabledModels []database.ChatModelConfig
|
||||
enabledModels []database.GetEnabledChatModelConfigsRow
|
||||
providerStatus map[string]chatprovider.ProviderAvailability
|
||||
providerStatusByID map[uuid.UUID]chatprovider.ProviderAvailability
|
||||
enabledProviderNames map[string]struct{}
|
||||
@@ -886,8 +886,8 @@ func (api *API) getUserChatProviderAvailability(
|
||||
if normalizedProvider == "" {
|
||||
continue
|
||||
}
|
||||
if model.AIProviderID.Valid {
|
||||
status, ok := availability.providerStatusByID[model.AIProviderID.UUID]
|
||||
if model.ChatModelConfig.AIProviderID.Valid {
|
||||
status, ok := availability.providerStatusByID[model.ChatModelConfig.AIProviderID.UUID]
|
||||
if ok {
|
||||
mergeProviderStatus(modelStatusByType, normalizedProvider, status)
|
||||
}
|
||||
@@ -904,8 +904,8 @@ func (api *API) getUserChatProviderAvailability(
|
||||
|
||||
for _, model := range enabledModels {
|
||||
normalizedProvider := chatprovider.NormalizeProvider(model.Provider)
|
||||
if model.AIProviderID.Valid {
|
||||
status, ok := availability.providerStatusByID[model.AIProviderID.UUID]
|
||||
if model.ChatModelConfig.AIProviderID.Valid {
|
||||
status, ok := availability.providerStatusByID[model.ChatModelConfig.AIProviderID.UUID]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
@@ -915,8 +915,8 @@ func (api *API) getUserChatProviderAvailability(
|
||||
}
|
||||
availability.configuredModels = append(availability.configuredModels, chatprovider.ConfiguredModel{
|
||||
Provider: model.Provider,
|
||||
Model: model.Model,
|
||||
DisplayName: model.DisplayName,
|
||||
Model: model.ChatModelConfig.Model,
|
||||
DisplayName: model.ChatModelConfig.DisplayName,
|
||||
})
|
||||
}
|
||||
return availability, nil
|
||||
@@ -966,21 +966,10 @@ func (api *API) userCanUseChatModelConfig(
|
||||
}
|
||||
return chatModelConfigAvailable, nil
|
||||
}
|
||||
provider, _, err := chatprovider.ResolveModelWithProviderHint(model.Model, model.Provider)
|
||||
if err != nil {
|
||||
return chatModelConfigUnavailableProviderDisabled, nil
|
||||
}
|
||||
if _, ok := availability.enabledProviderNames[provider]; !ok {
|
||||
return chatModelConfigUnavailableProviderDisabled, nil
|
||||
}
|
||||
providerStatus, ok := availability.providerStatus[provider]
|
||||
if !ok {
|
||||
return chatModelConfigUnavailableProviderDisabled, nil
|
||||
}
|
||||
if !providerStatus.Available {
|
||||
return chatModelConfigUnavailableCredentialsMissing, nil
|
||||
}
|
||||
return chatModelConfigAvailable, nil
|
||||
// Active configs always carry a provider FK (CHECK
|
||||
// chat_model_configs_ai_provider_required_when_active), so an unset FK
|
||||
// means the config is not usable.
|
||||
return chatModelConfigUnavailableModelNotFoundOrDisabled, nil
|
||||
}
|
||||
|
||||
func (api *API) validateUserChatModelConfigAvailable(
|
||||
@@ -6828,7 +6817,12 @@ func (api *API) listChatModelConfigs(rw http.ResponseWriter, r *http.Request) {
|
||||
configs, err = api.Database.GetChatModelConfigs(ctx)
|
||||
} else {
|
||||
//nolint:gocritic // All authenticated users need to read enabled model configs to use the chat feature.
|
||||
configs, err = api.Database.GetEnabledChatModelConfigs(dbauthz.AsChatd(ctx))
|
||||
rows, rowsErr := api.Database.GetEnabledChatModelConfigs(dbauthz.AsChatd(ctx))
|
||||
err = rowsErr
|
||||
configs = make([]database.ChatModelConfig, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
configs = append(configs, row.ChatModelConfig)
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
@@ -6900,7 +6894,6 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
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)
|
||||
@@ -6956,7 +6949,6 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
insertParams := database.InsertChatModelConfigParams{
|
||||
Provider: provider,
|
||||
Model: model,
|
||||
DisplayName: strings.TrimSpace(req.DisplayName),
|
||||
Enabled: enabled,
|
||||
@@ -6982,7 +6974,6 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
if !lockedAIProvider.Enabled {
|
||||
return errChatProviderNotConfigured
|
||||
}
|
||||
insertParams.Provider = string(lockedAIProvider.Type)
|
||||
if err := validateChatModelConfigProviderModel(lockedAIProvider, insertParams.Model); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -7087,19 +7078,6 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
if strings.TrimSpace(req.Provider) != "" && req.AIProviderID == nil {
|
||||
requestedProvider := chatprovider.NormalizeProvider(req.Provider)
|
||||
if requestedProvider == "" {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{Message: "Invalid provider."})
|
||||
return
|
||||
}
|
||||
if requestedProvider != existing.Provider {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{Message: "AI provider ID is required when updating provider."})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
provider := existing.Provider
|
||||
aiProviderID := existing.AIProviderID
|
||||
if req.AIProviderID != nil {
|
||||
//nolint:gocritic // The route already authorized chat model config updates.
|
||||
@@ -7119,7 +7097,6 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
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}
|
||||
}
|
||||
|
||||
@@ -7179,7 +7156,6 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
updateParams := database.UpdateChatModelConfigParams{
|
||||
Provider: provider,
|
||||
Model: model,
|
||||
DisplayName: displayName,
|
||||
Enabled: enabled,
|
||||
@@ -7208,7 +7184,6 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) {
|
||||
if !aiProvider.Enabled {
|
||||
return errChatProviderNotConfigured
|
||||
}
|
||||
updateParams.Provider = string(aiProvider.Type)
|
||||
if err := validateChatModelConfigProviderModel(aiProvider, updateParams.Model); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -7388,7 +7363,6 @@ func chatModelConfigToUpdateParams(
|
||||
config database.ChatModelConfig,
|
||||
) database.UpdateChatModelConfigParams {
|
||||
return database.UpdateChatModelConfigParams{
|
||||
Provider: config.Provider,
|
||||
Model: config.Model,
|
||||
DisplayName: config.DisplayName,
|
||||
Enabled: config.Enabled,
|
||||
@@ -7458,14 +7432,11 @@ func parseChatModelConfigID(rw http.ResponseWriter, r *http.Request) (uuid.UUID,
|
||||
}
|
||||
|
||||
func convertChatModelConfig(config database.ChatModelConfig) codersdk.ChatModelConfig {
|
||||
var aiProviderID *uuid.UUID
|
||||
if config.AIProviderID.Valid {
|
||||
aiProviderID = &config.AIProviderID.UUID
|
||||
}
|
||||
// Active configs always carry a non-null ai_provider_id (CHECK
|
||||
// chat_model_configs_ai_provider_required_when_active).
|
||||
return codersdk.ChatModelConfig{
|
||||
ID: config.ID,
|
||||
Provider: config.Provider,
|
||||
AIProviderID: aiProviderID,
|
||||
AIProviderID: config.AIProviderID.UUID,
|
||||
Model: config.Model,
|
||||
DisplayName: config.DisplayName,
|
||||
Enabled: config.Enabled,
|
||||
|
||||
+33
-57
@@ -1696,7 +1696,7 @@ func TestListChatModels(t *testing.T) {
|
||||
|
||||
var openAIProvider *codersdk.ChatModelProvider
|
||||
for i := range models.Providers {
|
||||
if models.Providers[i].Provider == modelConfig.Provider {
|
||||
if models.Providers[i].Provider == coderdtest.TestChatProviderOpenAICompat {
|
||||
openAIProvider = &models.Providers[i]
|
||||
break
|
||||
}
|
||||
@@ -1706,7 +1706,7 @@ func TestListChatModels(t *testing.T) {
|
||||
|
||||
foundModel := false
|
||||
for _, model := range openAIProvider.Models {
|
||||
if model.Provider == modelConfig.Provider && model.Model == modelConfig.Model {
|
||||
if model.Provider == coderdtest.TestChatProviderOpenAICompat && model.Model == modelConfig.Model {
|
||||
foundModel = true
|
||||
break
|
||||
}
|
||||
@@ -1772,14 +1772,14 @@ func TestListChatModels(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
client := newChatClient(t)
|
||||
_ = coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
_ = createChatModelConfig(t, client)
|
||||
|
||||
models, err := client.ListChatModels(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
var openAIProvider *codersdk.ChatModelProvider
|
||||
for i := range models.Providers {
|
||||
if models.Providers[i].Provider == modelConfig.Provider {
|
||||
if models.Providers[i].Provider == coderdtest.TestChatProviderOpenAICompat {
|
||||
openAIProvider = &models.Providers[i]
|
||||
break
|
||||
}
|
||||
@@ -1800,7 +1800,6 @@ func TestListChatModels(t *testing.T) {
|
||||
|
||||
contextLimit := int64(4096)
|
||||
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: string(providerType),
|
||||
AIProviderID: &provider.ID,
|
||||
Model: "claude-sonnet",
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -1851,7 +1850,6 @@ func TestListChatModels(t *testing.T) {
|
||||
|
||||
contextLimit := int64(4096)
|
||||
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "google",
|
||||
AIProviderID: &provider.ID,
|
||||
Model: "gemini-1.5-pro",
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -1903,7 +1901,6 @@ func TestListChatModels(t *testing.T) {
|
||||
|
||||
contextLimit := int64(4096)
|
||||
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai",
|
||||
AIProviderID: &provider.ID,
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -2397,14 +2394,14 @@ func TestListChatProviders(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
client := newChatClient(t)
|
||||
_ = coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
_ = createChatModelConfig(t, client)
|
||||
|
||||
providers, err := client.ListChatProviders(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
var openAIProvider *codersdk.ChatProviderConfig
|
||||
for i := range providers {
|
||||
if providers[i].Provider == modelConfig.Provider {
|
||||
if providers[i].Provider == coderdtest.TestChatProviderOpenAICompat {
|
||||
openAIProvider = &providers[i]
|
||||
break
|
||||
}
|
||||
@@ -3542,7 +3539,7 @@ func TestListChatModelConfigs(t *testing.T) {
|
||||
for _, config := range configs {
|
||||
if config.ID == modelConfig.ID {
|
||||
found = true
|
||||
require.Equal(t, modelConfig.Provider, config.Provider)
|
||||
require.Equal(t, modelConfig.AIProviderID, config.AIProviderID)
|
||||
require.Equal(t, modelConfig.Model, config.Model)
|
||||
require.True(t, config.IsDefault)
|
||||
}
|
||||
@@ -3562,7 +3559,6 @@ func TestListChatModelConfigs(t *testing.T) {
|
||||
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",
|
||||
@@ -3599,8 +3595,7 @@ func TestListChatModelConfigs(t *testing.T) {
|
||||
contextLimit := int64(4096)
|
||||
enabled := false
|
||||
_, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: enabledConfig.Provider,
|
||||
AIProviderID: enabledConfig.AIProviderID,
|
||||
AIProviderID: &enabledConfig.AIProviderID,
|
||||
Model: "gpt-4o-disabled",
|
||||
DisplayName: "GPT-4o Disabled",
|
||||
Enabled: &enabled,
|
||||
@@ -3626,7 +3621,6 @@ func TestListChatModelConfigs(t *testing.T) {
|
||||
|
||||
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",
|
||||
@@ -3670,7 +3664,7 @@ func TestListChatModelConfigs(t *testing.T) {
|
||||
for _, config := range configs {
|
||||
if config.ID == modelConfig.ID {
|
||||
found = true
|
||||
require.Equal(t, modelConfig.Provider, config.Provider)
|
||||
require.Equal(t, modelConfig.AIProviderID, config.AIProviderID)
|
||||
require.Equal(t, modelConfig.Model, config.Model)
|
||||
}
|
||||
}
|
||||
@@ -3701,7 +3695,6 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
},
|
||||
}
|
||||
modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai",
|
||||
AIProviderID: &aiProvider.ID,
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -3710,7 +3703,7 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotEqual(t, uuid.Nil, modelConfig.ID)
|
||||
require.Equal(t, "openai", modelConfig.Provider)
|
||||
require.Equal(t, aiProvider.ID, modelConfig.AIProviderID)
|
||||
require.Equal(t, "gpt-4o-mini", modelConfig.Model)
|
||||
require.EqualValues(t, 4096, modelConfig.ContextLimit)
|
||||
require.True(t, modelConfig.IsDefault)
|
||||
@@ -3733,7 +3726,6 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
|
||||
contextLimit := int64(4096)
|
||||
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai",
|
||||
AIProviderID: &aiProvider.ID,
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -3761,7 +3753,6 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
aiProvider := createAIProviderForTest(t, client, "openai", "test-api-key")
|
||||
|
||||
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai",
|
||||
AIProviderID: &aiProvider.ID,
|
||||
Model: "gpt-4o-mini",
|
||||
})
|
||||
@@ -3778,7 +3769,6 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
|
||||
contextLimit := int64(4096)
|
||||
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
})
|
||||
@@ -3796,7 +3786,6 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
contextLimit := int64(4096)
|
||||
missingProviderID := uuid.New()
|
||||
_, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai",
|
||||
AIProviderID: &missingProviderID,
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -3826,9 +3815,8 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
ContextLimit: &contextLimit,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "openai", modelConfig.Provider)
|
||||
require.NotNil(t, modelConfig.AIProviderID)
|
||||
require.Equal(t, provider.ID, *modelConfig.AIProviderID)
|
||||
require.NotEqual(t, uuid.Nil, modelConfig.AIProviderID)
|
||||
require.Equal(t, provider.ID, modelConfig.AIProviderID)
|
||||
})
|
||||
|
||||
t.Run("AIProviderIDNotConfigured", func(t *testing.T) {
|
||||
@@ -3913,7 +3901,6 @@ func TestCreateChatModelConfig(t *testing.T) {
|
||||
|
||||
contextLimit := int64(4096)
|
||||
_, err := memberClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "openai",
|
||||
AIProviderID: &aiProvider.ID,
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -3968,14 +3955,12 @@ func TestUpdateChatModelConfig(t *testing.T) {
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
updated, err := client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
|
||||
Provider: modelConfig.Provider,
|
||||
Model: "gpt-4o-mini-updated",
|
||||
Model: "gpt-4o-mini-updated",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, modelConfig.ID, updated.ID)
|
||||
require.Equal(t, modelConfig.Provider, updated.Provider)
|
||||
require.NotNil(t, updated.AIProviderID)
|
||||
require.Equal(t, *modelConfig.AIProviderID, *updated.AIProviderID)
|
||||
require.NotEqual(t, uuid.Nil, updated.AIProviderID)
|
||||
require.Equal(t, modelConfig.AIProviderID, updated.AIProviderID)
|
||||
require.Equal(t, "gpt-4o-mini-updated", updated.Model)
|
||||
})
|
||||
|
||||
@@ -4028,7 +4013,6 @@ func TestUpdateChatModelConfig(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: string(database.AIProviderTypeOpenai),
|
||||
Model: "anthropic/claude-opus-4.6",
|
||||
AIProviderID: uuid.NullUUID{UUID: aiProvider.ID, Valid: true},
|
||||
})
|
||||
@@ -4132,7 +4116,6 @@ func TestUpdateChatModelConfig(t *testing.T) {
|
||||
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",
|
||||
@@ -4218,9 +4201,8 @@ func TestUpdateChatModelConfig(t *testing.T) {
|
||||
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)
|
||||
require.NotEqual(t, uuid.Nil, updated.AIProviderID)
|
||||
require.Equal(t, provider.ID, updated.AIProviderID)
|
||||
})
|
||||
|
||||
t.Run("UpdateProviderPreservesAIProviderIDWhenTypeUnchanged", func(t *testing.T) {
|
||||
@@ -4243,15 +4225,14 @@ func TestUpdateChatModelConfig(t *testing.T) {
|
||||
Model: "claude-3-5-sonnet-latest",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, updated.AIProviderID)
|
||||
require.NotEqual(t, uuid.Nil, updated.AIProviderID)
|
||||
|
||||
updated, err = client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{
|
||||
Provider: "anthropic",
|
||||
Model: "claude-3-5-haiku-latest",
|
||||
Model: "claude-3-5-haiku-latest",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, updated.AIProviderID)
|
||||
require.Equal(t, provider.ID, *updated.AIProviderID)
|
||||
require.NotEqual(t, uuid.Nil, updated.AIProviderID)
|
||||
require.Equal(t, provider.ID, updated.AIProviderID)
|
||||
})
|
||||
|
||||
t.Run("UpdateAIProviderIDNotConfigured", func(t *testing.T) {
|
||||
@@ -4354,7 +4335,6 @@ func TestUpdateChatModelConfig(t *testing.T) {
|
||||
contextLimit := int64(4096)
|
||||
isDefault := false
|
||||
candidateConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: &aiProvider.ID,
|
||||
Model: "claude-3-5-sonnet",
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -6635,7 +6615,7 @@ func TestSendMessageWithModelOverrideUpdatesLastModelConfigID(t *testing.T) {
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfigA := createChatModelConfig(t, client)
|
||||
modelConfigB := createAdditionalChatModelConfig(t, client, modelConfigA.Provider, "gpt-4o-mini-override-"+uuid.NewString())
|
||||
modelConfigB := createAdditionalChatModelConfig(t, client, coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-override-"+uuid.NewString())
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: user.OrganizationID,
|
||||
@@ -6679,7 +6659,7 @@ func TestSendMessageQueuesEffectiveModelConfigID(t *testing.T) {
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfigA := createChatModelConfig(t, client)
|
||||
modelConfigB := createAdditionalChatModelConfig(t, client, modelConfigA.Provider, "gpt-4o-mini-queued-"+uuid.NewString())
|
||||
modelConfigB := createAdditionalChatModelConfig(t, client, coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-queued-"+uuid.NewString())
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: user.OrganizationID,
|
||||
@@ -6730,7 +6710,7 @@ func TestQueuedMessageWithoutOverrideCapturesEnqueueTimeModel(t *testing.T) {
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfigA := createChatModelConfig(t, client)
|
||||
modelConfigB := createAdditionalChatModelConfig(t, client, modelConfigA.Provider, "gpt-4o-mini-later-"+uuid.NewString())
|
||||
modelConfigB := createAdditionalChatModelConfig(t, client, coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-later-"+uuid.NewString())
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: user.OrganizationID,
|
||||
@@ -6824,7 +6804,7 @@ func TestWatchChatsStatusChangeCarriesUpdatedLastModelConfigID(t *testing.T) {
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfigA := createChatModelConfig(t, client)
|
||||
modelConfigB := createAdditionalChatModelConfig(t, client, modelConfigA.Provider, "gpt-4o-mini-watch-direct-"+uuid.NewString())
|
||||
modelConfigB := createAdditionalChatModelConfig(t, client, coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-watch-direct-"+uuid.NewString())
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: user.OrganizationID,
|
||||
@@ -6857,7 +6837,7 @@ func TestWatchChatsStatusChangeCarriesUpdatedLastModelConfigID(t *testing.T) {
|
||||
client, db := newChatClientWithDatabase(t)
|
||||
user := coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfigA := createChatModelConfig(t, client)
|
||||
modelConfigB := createAdditionalChatModelConfig(t, client, modelConfigA.Provider, "gpt-4o-mini-watch-promote-"+uuid.NewString())
|
||||
modelConfigB := createAdditionalChatModelConfig(t, client, coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-watch-promote-"+uuid.NewString())
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: user.OrganizationID,
|
||||
@@ -8177,7 +8157,7 @@ func TestPatchChatMessage(t *testing.T) {
|
||||
overrideModel := createAdditionalChatModelConfig(
|
||||
t,
|
||||
client,
|
||||
defaultModel.Provider,
|
||||
coderdtest.TestChatProviderOpenAICompat,
|
||||
"gpt-4o-mini-edit-override",
|
||||
)
|
||||
|
||||
@@ -10910,7 +10890,6 @@ func createAdditionalChatModelConfig(
|
||||
contextLimit := int64(4096)
|
||||
isDefault := false
|
||||
modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: provider,
|
||||
AIProviderID: &aiProvider.ID,
|
||||
Model: model,
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -11566,13 +11545,13 @@ func TestChatModelOverrides(t *testing.T) {
|
||||
openAIModel := createAdditionalChatModelConfig(
|
||||
t,
|
||||
adminClient,
|
||||
defaultModel.Provider,
|
||||
coderdtest.TestChatProviderOpenAICompat,
|
||||
"gpt-4.1-mini-"+string(setting.context),
|
||||
)
|
||||
disabledModel := createDisabledChatModelConfig(
|
||||
t,
|
||||
adminClient,
|
||||
defaultModel.Provider,
|
||||
coderdtest.TestChatProviderOpenAICompat,
|
||||
"gpt-4.1-disabled-"+string(setting.context),
|
||||
)
|
||||
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
|
||||
@@ -11784,7 +11763,7 @@ func TestUserChatPersonalModelOverrides(t *testing.T) {
|
||||
noKeyClient := codersdk.NewExperimentalClient(noKeyClientRaw)
|
||||
|
||||
defaultModelConfig := createChatModelConfig(t, adminClient)
|
||||
provider := enableUserChatProviderKey(t, adminClient, memberClient, defaultModelConfig.Provider)
|
||||
provider := enableUserChatProviderKey(t, adminClient, memberClient, coderdtest.TestChatProviderOpenAICompat)
|
||||
modelProvider := createAIProviderForTest(t, adminClient, "anthropic", "")
|
||||
_, err := memberClient.UpsertUserAIProviderKey(ctx, "me", modelProvider.ID, codersdk.CreateUserAIProviderKeyRequest{
|
||||
APIKey: "test-user-api-key-" + uuid.NewString(),
|
||||
@@ -11792,7 +11771,6 @@ func TestUserChatPersonalModelOverrides(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
contextLimit := int64(4096)
|
||||
modelConfig, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: &modelProvider.ID,
|
||||
Model: "claude-personal-" + uuid.NewString(),
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -11810,13 +11788,12 @@ func TestUserChatPersonalModelOverrides(t *testing.T) {
|
||||
disabledModelConfig := createDisabledChatModelConfig(
|
||||
t,
|
||||
adminClient,
|
||||
defaultModelConfig.Provider,
|
||||
coderdtest.TestChatProviderOpenAICompat,
|
||||
"gpt-4o-personal-disabled-"+uuid.NewString(),
|
||||
)
|
||||
disabledProvider := createAIProviderForTest(t, adminClient, "google", "test-api-key")
|
||||
contextLimit = int64(4096)
|
||||
disabledProviderModelConfig, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "google",
|
||||
AIProviderID: &disabledProvider.ID,
|
||||
Model: "gemini-personal-disabled-provider-" + uuid.NewString(),
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -12153,7 +12130,7 @@ func TestCreateChatPersonalModelOverrideRoot(t *testing.T) {
|
||||
adminClient, db := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, adminClient.Client)
|
||||
defaultModel := createChatModelConfig(t, adminClient)
|
||||
_ = enableUserChatProviderKey(t, adminClient, adminClient, defaultModel.Provider)
|
||||
_ = enableUserChatProviderKey(t, adminClient, adminClient, coderdtest.TestChatProviderOpenAICompat)
|
||||
overrideProvider := createAIProviderForTest(t, adminClient, "anthropic", "")
|
||||
_, err := adminClient.UpsertUserAIProviderKey(ctx, "me", overrideProvider.ID, codersdk.CreateUserAIProviderKeyRequest{
|
||||
APIKey: "test-user-api-key-" + uuid.NewString(),
|
||||
@@ -12161,7 +12138,6 @@ func TestCreateChatPersonalModelOverrideRoot(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
contextLimit := int64(4096)
|
||||
overrideModel, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: &overrideProvider.ID,
|
||||
Model: "claude-root-personal-" + uuid.NewString(),
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -12170,7 +12146,7 @@ func TestCreateChatPersonalModelOverrideRoot(t *testing.T) {
|
||||
disabledModel := createDisabledChatModelConfig(
|
||||
t,
|
||||
adminClient,
|
||||
defaultModel.Provider,
|
||||
coderdtest.TestChatProviderOpenAICompat,
|
||||
"gpt-4o-root-personal-disabled-"+uuid.NewString(),
|
||||
)
|
||||
memberClientRaw, member := coderdtest.CreateAnotherUser(
|
||||
|
||||
@@ -1600,18 +1600,18 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
|
||||
// Create chat providers (required FK for model configs).
|
||||
_ = dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
anthropicProvider := dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
Provider: "anthropic",
|
||||
DisplayName: "Anthropic",
|
||||
})
|
||||
_ = dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
openaiProvider := dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
Provider: "openai",
|
||||
DisplayName: "OpenAI",
|
||||
})
|
||||
|
||||
// Create a model config.
|
||||
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: uuid.NullUUID{UUID: anthropicProvider.ID, Valid: true},
|
||||
Model: "claude-sonnet-4-20250514",
|
||||
DisplayName: "Claude Sonnet",
|
||||
IsDefault: true,
|
||||
@@ -1620,14 +1620,14 @@ func TestChatsTelemetry(t *testing.T) {
|
||||
|
||||
// Create a second model config to test full dump.
|
||||
modelCfg2 := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o",
|
||||
DisplayName: "GPT-4o",
|
||||
AIProviderID: uuid.NullUUID{UUID: openaiProvider.ID, Valid: true},
|
||||
Model: "gpt-4o",
|
||||
DisplayName: "GPT-4o",
|
||||
})
|
||||
|
||||
// Create a soft-deleted model config — should NOT appear in telemetry.
|
||||
deletedCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: uuid.NullUUID{UUID: anthropicProvider.ID, Valid: true},
|
||||
Model: "claude-deleted",
|
||||
DisplayName: "Deleted Model",
|
||||
ContextLimit: 100000,
|
||||
@@ -1948,13 +1948,13 @@ func TestChatDiffStatusSummaryTelemetry(t *testing.T) {
|
||||
org, err := db.GetDefaultOrganization(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
_ = dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
anthropicProvider := dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
Provider: "anthropic",
|
||||
DisplayName: "Anthropic",
|
||||
})
|
||||
|
||||
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: uuid.NullUUID{UUID: anthropicProvider.ID, Valid: true},
|
||||
Model: "claude-sonnet-4-20250514",
|
||||
DisplayName: "Claude Sonnet",
|
||||
IsDefault: true,
|
||||
|
||||
@@ -118,7 +118,6 @@ func insertAgentChatTestModelConfig(
|
||||
})
|
||||
|
||||
return dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
CreatedBy: createdBy,
|
||||
UpdatedBy: createdBy,
|
||||
|
||||
@@ -248,7 +248,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
|
||||
getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) {
|
||||
return database.ChatModelConfig{
|
||||
ID: configID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-5.2",
|
||||
Enabled: true,
|
||||
CreatedAt: time.Unix(0, 0).UTC(),
|
||||
@@ -283,7 +282,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
|
||||
getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) {
|
||||
return database.ChatModelConfig{
|
||||
ID: configID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-5.2",
|
||||
Enabled: true,
|
||||
CreatedAt: time.Unix(0, 0).UTC(),
|
||||
@@ -322,6 +320,7 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
configID := uuid.New()
|
||||
providerID := uuid.New()
|
||||
rawOptions, err := json.Marshal(codersdk.ChatModelCallConfig{
|
||||
Temperature: func() *float64 { v := 0.42; return &v }(),
|
||||
})
|
||||
@@ -329,16 +328,29 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
|
||||
store := &advisorOverrideStubStore{
|
||||
getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) {
|
||||
return database.ChatModelConfig{
|
||||
ID: configID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-5.2",
|
||||
Enabled: true,
|
||||
CreatedAt: time.Unix(0, 0).UTC(),
|
||||
UpdatedAt: time.Unix(0, 0).UTC(),
|
||||
Options: rawOptions,
|
||||
DisplayName: "gpt-5.2",
|
||||
ID: configID,
|
||||
Model: "gpt-5.2",
|
||||
Enabled: true,
|
||||
CreatedAt: time.Unix(0, 0).UTC(),
|
||||
UpdatedAt: time.Unix(0, 0).UTC(),
|
||||
Options: rawOptions,
|
||||
DisplayName: "gpt-5.2",
|
||||
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
|
||||
}, nil
|
||||
},
|
||||
getAIProviderByID: func(context.Context, uuid.UUID) (database.AIProvider, error) {
|
||||
return database.AIProvider{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
Enabled: true,
|
||||
}, nil
|
||||
},
|
||||
getAIProviderKeysByProviderID: func(context.Context, uuid.UUID) ([]database.AIProviderKey, error) {
|
||||
return []database.AIProviderKey{{
|
||||
ProviderID: providerID,
|
||||
APIKey: "sk-test",
|
||||
}}, nil
|
||||
},
|
||||
}
|
||||
p := newAdvisorTestServer(ctx, t, store)
|
||||
|
||||
@@ -373,7 +385,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
|
||||
getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) {
|
||||
return database.ChatModelConfig{
|
||||
ID: configID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-5.2",
|
||||
Enabled: true,
|
||||
CreatedAt: time.Unix(0, 0).UTC(),
|
||||
@@ -426,7 +437,6 @@ func TestResolveAdvisorModelOverridePromotesAIBridgeErrors(t *testing.T) {
|
||||
getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) {
|
||||
return database.ChatModelConfig{
|
||||
ID: configID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-5.2",
|
||||
Enabled: true,
|
||||
DisplayName: "gpt-5.2",
|
||||
|
||||
+12
-7
@@ -2588,6 +2588,16 @@ func (p *Server) prepareManualTitleDebugRun(
|
||||
finishDebugRun := func(error) {}
|
||||
|
||||
route, routeErr := p.resolveModelRouteForConfig(ctx, chat.OwnerID, modelConfig, keys)
|
||||
var routeProvider string
|
||||
if routeErr == nil {
|
||||
routeProvider, _ = route.providerHint()
|
||||
} else if modelConfig.AIProviderID.Valid {
|
||||
// Route resolution failed, but the linked provider still identifies the
|
||||
// type for the debug run record. Best-effort: leave empty if disabled.
|
||||
if provider, err := p.enabledAIProviderByID(ctx, modelConfig.AIProviderID.UUID); err == nil {
|
||||
routeProvider = string(provider.Type)
|
||||
}
|
||||
}
|
||||
debugOpts := modelOpts
|
||||
debugOpts.RecordHTTP = true
|
||||
var debugModelErr error
|
||||
@@ -2606,21 +2616,19 @@ func (p *Server) prepareManualTitleDebugRun(
|
||||
case debugModelErr != nil:
|
||||
p.logger.Warn(ctx, "failed to create debug-aware manual title model",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.F("provider", modelConfig.Provider),
|
||||
slog.F("model", modelConfig.Model),
|
||||
slog.Error(debugModelErr),
|
||||
)
|
||||
case debugModel == nil:
|
||||
p.logger.Warn(ctx, "manual title debug model creation returned nil",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.F("provider", modelConfig.Provider),
|
||||
slog.F("model", modelConfig.Model),
|
||||
)
|
||||
default:
|
||||
titleModel = chatdebug.WrapModel(debugModel, debugSvc, chatdebug.RecorderOptions{
|
||||
ChatID: chat.ID,
|
||||
OwnerID: chat.OwnerID,
|
||||
Provider: modelConfig.Provider,
|
||||
Provider: routeProvider,
|
||||
Model: modelConfig.Model,
|
||||
})
|
||||
}
|
||||
@@ -2651,7 +2659,7 @@ func (p *Server) prepareManualTitleDebugRun(
|
||||
debugRun, createRunErr := debugSvc.CreateRun(createRunCtx, chatdebug.CreateRunParams{
|
||||
ChatID: chat.ID,
|
||||
ModelConfigID: modelConfig.ID,
|
||||
Provider: modelConfig.Provider,
|
||||
Provider: routeProvider,
|
||||
Model: modelConfig.Model,
|
||||
Kind: chatdebug.KindTitleGeneration,
|
||||
Status: chatdebug.StatusInProgress,
|
||||
@@ -2663,7 +2671,6 @@ func (p *Server) prepareManualTitleDebugRun(
|
||||
if createRunErr != nil {
|
||||
p.logger.Warn(ctx, "failed to create manual title debug run",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.F("provider", modelConfig.Provider),
|
||||
slog.F("model", modelConfig.Model),
|
||||
slog.Error(createRunErr),
|
||||
)
|
||||
@@ -2789,7 +2796,6 @@ func (p *Server) resolveManualTitleModel(
|
||||
if err != nil {
|
||||
p.logger.Debug(ctx, "manual title preferred model unavailable",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.F("provider", config.Provider),
|
||||
slog.F("model", config.Model),
|
||||
slog.Error(err),
|
||||
)
|
||||
@@ -2804,7 +2810,6 @@ func (p *Server) resolveManualTitleModel(
|
||||
if err != nil {
|
||||
p.logger.Debug(ctx, "manual title preferred model unavailable",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.F("provider", config.Provider),
|
||||
slog.F("model", config.Model),
|
||||
slog.Error(err),
|
||||
)
|
||||
|
||||
@@ -326,7 +326,6 @@ func seedAnthropicChatDependencies(t *testing.T, db database.Store, baseURL stri
|
||||
})
|
||||
dbgen.AIProviderKey(t, db, database.AIProviderKey{ProviderID: provider.ID})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "anthropic",
|
||||
Model: "claude-sonnet-4-20250514",
|
||||
IsDefault: true,
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
@@ -453,7 +452,6 @@ func updateModelForChainMode(t *testing.T, db database.Store, model database.Cha
|
||||
ID: model.ID,
|
||||
DisplayName: model.DisplayName,
|
||||
Model: model.Model,
|
||||
Provider: model.Provider,
|
||||
Enabled: model.Enabled,
|
||||
ContextLimit: model.ContextLimit,
|
||||
CompressionThreshold: model.CompressionThreshold,
|
||||
|
||||
@@ -26,6 +26,7 @@ import (
|
||||
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/rbac"
|
||||
"github.com/coder/coder/v2/coderd/workspacestats"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatdebug"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatloop"
|
||||
openaicomputeruse "github.com/coder/coder/v2/coderd/x/chatd/chatopenai/computeruse"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
@@ -794,11 +795,12 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) {
|
||||
WorkerID: uuid.NullUUID{UUID: workerID, Valid: true},
|
||||
Title: fallbackChatTitle(userPrompt),
|
||||
}
|
||||
providerID := uuid.New()
|
||||
modelConfig := database.ChatModelConfig{
|
||||
ID: modelConfigID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: 8192,
|
||||
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
|
||||
}
|
||||
updatedChat := chat
|
||||
updatedChat.Title = wantTitle
|
||||
@@ -832,14 +834,20 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) {
|
||||
}
|
||||
|
||||
db.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(modelConfig, nil)
|
||||
providerID := uuid.New()
|
||||
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(database.AIProvider{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
Enabled: true,
|
||||
BaseUrl: serverURL,
|
||||
}, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
Enabled: true,
|
||||
BaseUrl: serverURL,
|
||||
}}, nil)
|
||||
db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), []uuid.UUID{providerID}).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil)
|
||||
}}, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes()
|
||||
db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows)
|
||||
db.EXPECT().GetChatMessagesByChatIDAscPaginated(
|
||||
gomock.Any(),
|
||||
@@ -957,11 +965,12 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts_IdleChatReleasesManualLock(t
|
||||
lockedChat := chat
|
||||
lockedChat.WorkerID = uuid.NullUUID{UUID: manualTitleLockWorkerID, Valid: true}
|
||||
lockedChat.StartedAt = sql.NullTime{Time: time.Now(), Valid: true}
|
||||
providerID := uuid.New()
|
||||
modelConfig := database.ChatModelConfig{
|
||||
ID: modelConfigID,
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o-mini",
|
||||
ContextLimit: 8192,
|
||||
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
|
||||
}
|
||||
updatedChat := lockedChat
|
||||
updatedChat.Title = wantTitle
|
||||
@@ -998,14 +1007,20 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts_IdleChatReleasesManualLock(t
|
||||
}
|
||||
|
||||
db.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(modelConfig, nil)
|
||||
providerID := uuid.New()
|
||||
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(database.AIProvider{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
Enabled: true,
|
||||
BaseUrl: serverURL,
|
||||
}, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
Enabled: true,
|
||||
BaseUrl: serverURL,
|
||||
}}, nil)
|
||||
db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), []uuid.UUID{providerID}).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil)
|
||||
}}, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes()
|
||||
db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows)
|
||||
db.EXPECT().GetChatMessagesByChatIDAscPaginated(
|
||||
gomock.Any(),
|
||||
@@ -3495,3 +3510,73 @@ func TestServer_inflightContext(t *testing.T) {
|
||||
t.Fatal("inflight context not canceled on server shutdown")
|
||||
}
|
||||
}
|
||||
|
||||
// TestPrepareManualTitleDebugRun_RouteFailureDerivesProviderFromConfig drives
|
||||
// the fallback branch in prepareManualTitleDebugRun: AI-gateway route
|
||||
// resolution fails (the BYOK key lookup returns a non-ErrNoRows error) while
|
||||
// the linked provider stays enabled, so the debug run records the provider
|
||||
// type derived from modelConfig.AIProviderID instead of an empty string.
|
||||
func TestPrepareManualTitleDebugRun_RouteFailureDerivesProviderFromConfig(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
ctrl := gomock.NewController(t)
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
|
||||
ownerID := uuid.New()
|
||||
providerID := uuid.New()
|
||||
chat := database.Chat{ID: uuid.New(), OwnerID: ownerID}
|
||||
modelConfig := database.ChatModelConfig{
|
||||
ID: uuid.New(),
|
||||
Model: "claude-sonnet-4",
|
||||
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
|
||||
}
|
||||
provider := database.AIProvider{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeAnthropic,
|
||||
Name: "anthropic",
|
||||
Enabled: true,
|
||||
}
|
||||
|
||||
// Resolved twice: once by gatewayProviderForConfig during route resolution,
|
||||
// once by the fallback's own enabledAIProviderByID lookup.
|
||||
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
|
||||
// A non-ErrNoRows BYOK error fails route resolution while the provider stays
|
||||
// enabled, which is exactly the gap the fallback covers.
|
||||
db.EXPECT().GetUserAIProviderKeyByProviderID(gomock.Any(), database.GetUserAIProviderKeyByProviderIDParams{
|
||||
UserID: ownerID,
|
||||
AIProviderID: providerID,
|
||||
}).Return(database.UserAIProviderKey{}, sql.ErrConnDone)
|
||||
|
||||
var gotProvider sql.NullString
|
||||
db.EXPECT().InsertChatDebugRun(gomock.Any(), gomock.Any()).DoAndReturn(
|
||||
func(_ context.Context, params database.InsertChatDebugRunParams) (database.ChatDebugRun, error) {
|
||||
gotProvider = params.Provider
|
||||
return database.ChatDebugRun{ChatID: params.ChatID, Provider: params.Provider}, nil
|
||||
},
|
||||
)
|
||||
|
||||
server := &Server{
|
||||
db: db,
|
||||
logger: logger,
|
||||
aiGatewayRoutingEnabled: true,
|
||||
allowBYOK: true,
|
||||
}
|
||||
debugSvc := chatdebug.NewService(db, logger, nil)
|
||||
fallbackModel := &chattest.FakeModel{ProviderName: "stub", ModelName: "stub"}
|
||||
|
||||
server.prepareManualTitleDebugRun(
|
||||
ctx,
|
||||
debugSvc,
|
||||
chat,
|
||||
modelConfig,
|
||||
chatprovider.ProviderAPIKeys{},
|
||||
modelBuildOptions{},
|
||||
nil,
|
||||
fallbackModel,
|
||||
)
|
||||
|
||||
require.True(t, gotProvider.Valid, "debug run provider should be populated from the linked config")
|
||||
require.Equal(t, "anthropic", gotProvider.String)
|
||||
}
|
||||
|
||||
@@ -103,7 +103,6 @@ func TestActiveServer_RetryStreamSilenceTimeoutAndClassification(t *testing.T) {
|
||||
})
|
||||
user, org, _ := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o",
|
||||
Enabled: true,
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
|
||||
@@ -7713,7 +7713,6 @@ func updateChatModelCompressionThreshold(t *testing.T, db database.Store, model
|
||||
ID: model.ID,
|
||||
DisplayName: model.DisplayName,
|
||||
Model: model.Model,
|
||||
Provider: model.Provider,
|
||||
Enabled: model.Enabled,
|
||||
ContextLimit: model.ContextLimit,
|
||||
CompressionThreshold: model.CompressionThreshold,
|
||||
@@ -7730,7 +7729,6 @@ func updateChatModelContextLimit(t *testing.T, db database.Store, model database
|
||||
ID: model.ID,
|
||||
DisplayName: model.DisplayName,
|
||||
Model: model.Model,
|
||||
Provider: model.Provider,
|
||||
Enabled: model.Enabled,
|
||||
ContextLimit: model.ContextLimit,
|
||||
CompressionThreshold: model.CompressionThreshold,
|
||||
@@ -7749,7 +7747,6 @@ func updateChatModelCallConfig(t *testing.T, db database.Store, model database.C
|
||||
ID: model.ID,
|
||||
DisplayName: model.DisplayName,
|
||||
Model: model.Model,
|
||||
Provider: model.Provider,
|
||||
Enabled: model.Enabled,
|
||||
ContextLimit: model.ContextLimit,
|
||||
CompressionThreshold: model.CompressionThreshold,
|
||||
@@ -8333,6 +8330,8 @@ func TestProposeChatTitle_DebugRun(t *testing.T) {
|
||||
if tt.wantTitleGenerationRuns > 0 {
|
||||
require.Equal(t, string(codersdk.ChatDebugRunKindTitleGeneration), runs[0].Kind)
|
||||
require.Equal(t, string(tt.wantDebugStatus), runs[0].Status)
|
||||
require.True(t, runs[0].Provider.Valid)
|
||||
require.Equal(t, "openai", runs[0].Provider.String)
|
||||
require.True(t, runs[0].FinishedAt.Valid)
|
||||
require.True(t, runs[0].HistoryTipMessageID.Valid)
|
||||
require.Equal(t, message.ID, runs[0].HistoryTipMessageID.Int64)
|
||||
@@ -8378,14 +8377,14 @@ func seedChatDependenciesWithProvider(
|
||||
UserID: user.ID,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
providerConfig := dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
Provider: provider,
|
||||
DisplayName: provider,
|
||||
BaseUrl: baseURL,
|
||||
})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: provider,
|
||||
IsDefault: true,
|
||||
AIProviderID: uuid.NullUUID{UUID: providerConfig.ID, Valid: true},
|
||||
IsDefault: true,
|
||||
})
|
||||
return user, org, model
|
||||
}
|
||||
@@ -8423,8 +8422,8 @@ func seedChatDependenciesWithProviderPolicy(
|
||||
})
|
||||
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: provider,
|
||||
IsDefault: true,
|
||||
AIProviderID: uuid.NullUUID{UUID: providerConfig.ID, Valid: true},
|
||||
IsDefault: true,
|
||||
})
|
||||
|
||||
return user, org, providerConfig, model
|
||||
@@ -8482,13 +8481,32 @@ func insertChatModelConfigWithCallConfig(
|
||||
options, err := json.Marshal(callConfig)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Reuse the newest AI provider of this type (creating a bare one only when
|
||||
// none exists) so the config links the seeded provider carrying the mock
|
||||
// base URL and API key rather than a fresh credential-less one.
|
||||
providers, err := db.GetAIProviders(context.Background(), database.GetAIProvidersParams{IncludeDisabled: true})
|
||||
require.NoError(t, err)
|
||||
var aiProvider database.AIProvider
|
||||
for _, candidate := range providers {
|
||||
if candidate.Type != database.AIProviderType(provider) {
|
||||
continue
|
||||
}
|
||||
if aiProvider.ID == uuid.Nil || candidate.CreatedAt.After(aiProvider.CreatedAt) {
|
||||
aiProvider = candidate
|
||||
}
|
||||
}
|
||||
if aiProvider.ID == uuid.Nil {
|
||||
aiProvider = dbgen.AIProvider(t, db, database.AIProvider{
|
||||
Type: database.AIProviderType(provider),
|
||||
})
|
||||
}
|
||||
return dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: provider,
|
||||
Model: model,
|
||||
DisplayName: model,
|
||||
CreatedBy: uuid.NullUUID{UUID: userID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: userID, Valid: true},
|
||||
Options: options,
|
||||
AIProviderID: uuid.NullUUID{UUID: aiProvider.ID, Valid: true},
|
||||
Model: model,
|
||||
DisplayName: model,
|
||||
CreatedBy: uuid.NullUUID{UUID: userID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: userID, Valid: true},
|
||||
Options: options,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -9671,7 +9689,6 @@ func seedAIGatewayOpenAITestDependencies(
|
||||
BaseUrl: openAIURL,
|
||||
})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: string(database.AIProviderTypeOpenai),
|
||||
Model: "gpt-4o-mini",
|
||||
IsDefault: true,
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
@@ -12548,14 +12565,12 @@ func TestProviderSwitchSanitizesAndRestoresPEToolHistory(t *testing.T) {
|
||||
})
|
||||
|
||||
mA := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai-compat",
|
||||
Model: "gpt-4o-mini",
|
||||
DisplayName: "Model A",
|
||||
Enabled: true,
|
||||
AIProviderID: uuid.NullUUID{UUID: cpA.ID, Valid: true},
|
||||
})
|
||||
mB := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai-compat",
|
||||
Model: "gpt-4o-mini",
|
||||
DisplayName: "Model B",
|
||||
Enabled: true,
|
||||
|
||||
@@ -25,13 +25,14 @@ import (
|
||||
)
|
||||
|
||||
type testFixture struct {
|
||||
ctx context.Context
|
||||
db database.Store
|
||||
svc *chatdebug.Service
|
||||
org database.Organization
|
||||
owner database.User
|
||||
chat database.Chat
|
||||
model database.ChatModelConfig
|
||||
ctx context.Context
|
||||
db database.Store
|
||||
svc *chatdebug.Service
|
||||
org database.Organization
|
||||
owner database.User
|
||||
chat database.Chat
|
||||
model database.ChatModelConfig
|
||||
provider string
|
||||
}
|
||||
|
||||
func TestService_IsEnabled(t *testing.T) {
|
||||
@@ -115,7 +116,7 @@ func TestService_CreateRun(t *testing.T) {
|
||||
HistoryTipMessageID: historyTipMsg.ID,
|
||||
Kind: chatdebug.KindChatTurn,
|
||||
Status: chatdebug.StatusInProgress,
|
||||
Provider: fixture.model.Provider,
|
||||
Provider: fixture.provider,
|
||||
Model: fixture.model.Model,
|
||||
Summary: map[string]any{
|
||||
"phase": "create",
|
||||
@@ -126,7 +127,7 @@ func TestService_CreateRun(t *testing.T) {
|
||||
assertRunMatches(t, run, fixture.chat.ID, rootChat.ID, parentChat.ID,
|
||||
fixture.model.ID, triggerMsg.ID, historyTipMsg.ID,
|
||||
chatdebug.KindChatTurn, chatdebug.StatusInProgress,
|
||||
fixture.model.Provider, fixture.model.Model,
|
||||
fixture.provider, fixture.model.Model,
|
||||
`{"count":1,"phase":"create"}`)
|
||||
|
||||
stored, err := fixture.db.GetChatDebugRunByID(fixture.ctx, run.ID)
|
||||
@@ -470,7 +471,7 @@ func TestService_UpdateStep(t *testing.T) {
|
||||
ResponseStatus: 200,
|
||||
DurationMs: 25,
|
||||
}},
|
||||
Metadata: map[string]any{"provider": fixture.model.Provider},
|
||||
Metadata: map[string]any{"provider": fixture.provider},
|
||||
FinishedAt: finishedAt,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -487,7 +488,7 @@ func TestService_UpdateStep(t *testing.T) {
|
||||
`[{"number":1,"response_status":200,"duration_ms":25}]`,
|
||||
string(updated.Attempts),
|
||||
)
|
||||
require.JSONEq(t, `{"provider":"`+fixture.model.Provider+`"}`,
|
||||
require.JSONEq(t, `{"provider":"`+fixture.provider+`"}`,
|
||||
string(updated.Metadata))
|
||||
require.True(t, updated.FinishedAt.Valid)
|
||||
storedSteps, err := fixture.db.GetChatDebugStepsByRunID(fixture.ctx, run.ID)
|
||||
@@ -1070,14 +1071,17 @@ func newFixture(t *testing.T) testFixture {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
org, owner, chat, model := seedChat(t, db)
|
||||
prov, err := db.GetAIProviderByID(ctx, model.AIProviderID.UUID)
|
||||
require.NoError(t, err)
|
||||
return testFixture{
|
||||
ctx: ctx,
|
||||
db: db,
|
||||
svc: chatdebug.NewService(db, testutil.Logger(t), nil),
|
||||
org: org,
|
||||
owner: owner,
|
||||
chat: chat,
|
||||
model: model,
|
||||
ctx: ctx,
|
||||
db: db,
|
||||
svc: chatdebug.NewService(db, testutil.Logger(t), nil),
|
||||
org: org,
|
||||
owner: owner,
|
||||
chat: chat,
|
||||
model: model,
|
||||
provider: string(prov.Type),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1158,7 +1162,7 @@ func createRun(t *testing.T, fixture testFixture) database.ChatDebugRun {
|
||||
ModelConfigID: fixture.model.ID,
|
||||
Kind: chatdebug.KindChatTurn,
|
||||
Status: chatdebug.StatusInProgress,
|
||||
Provider: fixture.model.Provider,
|
||||
Provider: fixture.provider,
|
||||
Model: fixture.model.Model,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -2516,12 +2516,12 @@ func TestMediaToolResultRoundTrip(t *testing.T) {
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
|
||||
dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
anthropicProvider := dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
Provider: "anthropic",
|
||||
})
|
||||
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "anthropic",
|
||||
AIProviderID: uuid.NullUUID{UUID: anthropicProvider.ID, Valid: true},
|
||||
Model: "test-model",
|
||||
IsDefault: true,
|
||||
ContextLimit: 200000,
|
||||
|
||||
@@ -211,7 +211,6 @@ func seedFamilyDeps(t *testing.T, db database.Store) (database.User, database.Or
|
||||
BaseUrl: "http://example.invalid",
|
||||
})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
IsDefault: true,
|
||||
})
|
||||
return user, org, model
|
||||
|
||||
@@ -58,7 +58,6 @@ func newTestFixture(t *testing.T) *testFixture {
|
||||
BaseUrl: "http://example.invalid",
|
||||
})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
IsDefault: true,
|
||||
})
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
|
||||
@@ -38,7 +38,6 @@ func newTriggerFixture(t *testing.T) *triggerFixture {
|
||||
BaseUrl: "http://example.invalid",
|
||||
})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
IsDefault: true,
|
||||
})
|
||||
f := &testFixture{
|
||||
|
||||
@@ -683,7 +683,6 @@ func testAIProvider(name string) database.AIProvider {
|
||||
func testChatModelConfig(id uuid.UUID, model string) database.ChatModelConfig {
|
||||
return database.ChatModelConfig{
|
||||
ID: id,
|
||||
Provider: "openai",
|
||||
Model: model,
|
||||
DisplayName: model,
|
||||
Enabled: true,
|
||||
|
||||
@@ -262,7 +262,11 @@ func (server *Server) prepareGeneration(
|
||||
acceptsFilePart := func(mediaType string) bool {
|
||||
return chatprovider.AcceptsFilePartMediaType(model.Provider(), model.Model(), mediaType)
|
||||
}
|
||||
prompt, err = chatprompt.ConvertMessagesWithFiles(ctx, promptRows, server.chatFileResolver(modelConfig.Provider), logger, acceptsFilePart)
|
||||
providerType, err := modelRoute.providerHint()
|
||||
if err != nil {
|
||||
return xerrors.Errorf("resolve provider type: %w", err)
|
||||
}
|
||||
prompt, err = chatprompt.ConvertMessagesWithFiles(ctx, promptRows, server.chatFileResolver(providerType), logger, acceptsFilePart)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("build chat prompt: %w", err)
|
||||
}
|
||||
|
||||
@@ -113,7 +113,6 @@ func TestDeriveFinalTurnRunResult(t *testing.T) {
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
})
|
||||
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o-mini",
|
||||
DisplayName: "gpt-4o-mini",
|
||||
Options: json.RawMessage(`{}`),
|
||||
@@ -233,7 +232,6 @@ func TestDeriveFinalTurnRunResult(t *testing.T) {
|
||||
// degraded path that still returns the re-derived text and IDs.
|
||||
provider := insertInternalAIProvider(t, db, database.AIProviderTypeOpenai, "provider-api-key", false)
|
||||
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o-mini",
|
||||
DisplayName: "gpt-4o-mini",
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
|
||||
@@ -133,7 +133,6 @@ func newWorkerTestFixture(t *testing.T) *workerTestFixture {
|
||||
BaseUrl: "http://example.invalid",
|
||||
})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
IsDefault: true,
|
||||
})
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
|
||||
@@ -98,7 +98,6 @@ func TestAnthropicWebSearchRoundTrip(t *testing.T) {
|
||||
contextLimit := int64(200000)
|
||||
isDefault := true
|
||||
_, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: string(provider.Type),
|
||||
AIProviderID: &provider.ID,
|
||||
Model: "claude-sonnet-4-20250514",
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -358,7 +357,6 @@ func TestOpenAIReasoningRoundTrip(t *testing.T) {
|
||||
isDefault := true
|
||||
reasoningSummary := "auto"
|
||||
_, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: string(provider.Type),
|
||||
AIProviderID: &provider.ID,
|
||||
Model: "o4-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
@@ -508,7 +506,6 @@ func TestOpenAIReasoningRoundTripStoreFalse(t *testing.T) {
|
||||
isDefault := true
|
||||
reasoningSummary := "auto"
|
||||
_, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{
|
||||
Provider: string(provider.Type),
|
||||
AIProviderID: &provider.ID,
|
||||
Model: "o4-mini",
|
||||
ContextLimit: &contextLimit,
|
||||
|
||||
@@ -2,6 +2,7 @@ package chatd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"net/http"
|
||||
|
||||
"charm.land/fantasy"
|
||||
@@ -83,7 +84,7 @@ func (p *Server) directProviderHintAndProviderForConfig(
|
||||
modelConfig database.ChatModelConfig,
|
||||
) (string, *database.AIProvider, error) {
|
||||
if !modelConfig.AIProviderID.Valid {
|
||||
return modelConfig.Provider, nil, nil
|
||||
return "", nil, sql.ErrNoRows
|
||||
}
|
||||
provider, err := p.enabledAIProviderByID(ctx, modelConfig.AIProviderID.UUID)
|
||||
if err != nil {
|
||||
|
||||
@@ -125,7 +125,6 @@ func TestResolveModelRouteForConfigPreservesBaseURL(t *testing.T) {
|
||||
|
||||
server := &Server{db: db}
|
||||
route, err := server.resolveModelRouteForConfig(ctx, ownerID, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
|
||||
}, chatprovider.ProviderAPIKeys{})
|
||||
require.NoError(t, err)
|
||||
@@ -238,7 +237,6 @@ func TestResolveModelRouteForConfigAIGatewayProviderAuth(t *testing.T) {
|
||||
modelConfig := database.ChatModelConfig{
|
||||
ID: uuid.New(),
|
||||
Model: "gpt-4",
|
||||
Provider: "openai",
|
||||
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
|
||||
}
|
||||
|
||||
|
||||
@@ -102,17 +102,17 @@ func (p *Server) preferredShortTextCandidates(
|
||||
}
|
||||
|
||||
func selectPreferredConfiguredShortTextModelConfig(
|
||||
configs []database.ChatModelConfig,
|
||||
configs []database.GetEnabledChatModelConfigsRow,
|
||||
) (database.ChatModelConfig, bool) {
|
||||
for _, preferred := range preferredTitleModels {
|
||||
for _, config := range configs {
|
||||
if chatprovider.NormalizeProvider(config.Provider) != preferred.provider {
|
||||
continue
|
||||
}
|
||||
if !strings.EqualFold(strings.TrimSpace(config.Model), preferred.model) {
|
||||
if !strings.EqualFold(strings.TrimSpace(config.ChatModelConfig.Model), preferred.model) {
|
||||
continue
|
||||
}
|
||||
return config, true
|
||||
return config.ChatModelConfig, true
|
||||
}
|
||||
}
|
||||
return database.ChatModelConfig{}, false
|
||||
@@ -181,11 +181,18 @@ func (p *Server) GenerateChatTitleAsync(ctx context.Context, chat database.Chat)
|
||||
)
|
||||
return
|
||||
}
|
||||
providerType, err := route.providerHint()
|
||||
if err != nil {
|
||||
logger.Debug(titleCtx, "failed to resolve provider type for automatic title generation",
|
||||
slog.Error(err),
|
||||
)
|
||||
return
|
||||
}
|
||||
p.maybeGenerateChatTitle(
|
||||
turnCtx,
|
||||
chat,
|
||||
messages,
|
||||
modelConfig.Provider,
|
||||
providerType,
|
||||
modelConfig.Model,
|
||||
model,
|
||||
route,
|
||||
@@ -259,8 +266,16 @@ func (p *Server) maybeGenerateChatTitle(
|
||||
|
||||
var candidates []shortTextCandidate
|
||||
if overrideSet {
|
||||
overrideProvider, err := overrideRoute.providerHint()
|
||||
if err != nil {
|
||||
logger.Debug(ctx, "failed to resolve provider type for title generation override",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.Error(err),
|
||||
)
|
||||
return
|
||||
}
|
||||
candidates = []shortTextCandidate{{
|
||||
provider: overrideConfig.Provider,
|
||||
provider: overrideProvider,
|
||||
model: overrideConfig.Model,
|
||||
route: overrideRoute,
|
||||
lm: overrideModel,
|
||||
|
||||
@@ -391,8 +391,7 @@ func TestMaybeGenerateChatTitlePreservesUpdatedAt(t *testing.T) {
|
||||
CentralApiKeyEnabled: true,
|
||||
})
|
||||
modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
Model: "test-model",
|
||||
})
|
||||
|
||||
userPrompt := "summarize failed workspace build logs"
|
||||
@@ -586,24 +585,23 @@ func Test_selectPreferredConfiguredShortTextModelConfig(t *testing.T) {
|
||||
t.Run("chooses the highest-priority configured lightweight model", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
configs := []database.ChatModelConfig{
|
||||
{Provider: preferredTitleModels[2].provider, Model: preferredTitleModels[2].model},
|
||||
{Provider: preferredTitleModels[1].provider, Model: preferredTitleModels[1].model},
|
||||
{Provider: "openai", Model: "gpt-4.1"},
|
||||
configs := []database.GetEnabledChatModelConfigsRow{
|
||||
{ChatModelConfig: database.ChatModelConfig{Model: preferredTitleModels[2].model}, Provider: preferredTitleModels[2].provider},
|
||||
{ChatModelConfig: database.ChatModelConfig{Model: preferredTitleModels[1].model}, Provider: preferredTitleModels[1].provider},
|
||||
{ChatModelConfig: database.ChatModelConfig{Model: "gpt-4.1"}, Provider: "openai"},
|
||||
}
|
||||
|
||||
got, ok := selectPreferredConfiguredShortTextModelConfig(configs)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, preferredTitleModels[1].provider, got.Provider)
|
||||
require.Equal(t, preferredTitleModels[1].model, got.Model)
|
||||
})
|
||||
|
||||
t.Run("returns false when no preferred lightweight model is configured", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got, ok := selectPreferredConfiguredShortTextModelConfig([]database.ChatModelConfig{{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4.1",
|
||||
got, ok := selectPreferredConfiguredShortTextModelConfig([]database.GetEnabledChatModelConfigsRow{{
|
||||
ChatModelConfig: database.ChatModelConfig{Model: "gpt-4.1"},
|
||||
Provider: "openai",
|
||||
}})
|
||||
require.False(t, ok)
|
||||
require.Equal(t, database.ChatModelConfig{}, got)
|
||||
|
||||
@@ -161,39 +161,6 @@ func personalModelOverrideContextForSubagent(
|
||||
}
|
||||
}
|
||||
|
||||
func validateModelConfigAndResolveProvider(
|
||||
modelConfig database.ChatModelConfig,
|
||||
) (database.ChatModelConfig, string, error) {
|
||||
if !modelConfig.Enabled {
|
||||
return database.ChatModelConfig{}, "", sql.ErrNoRows
|
||||
}
|
||||
providerName, _, err := chatprovider.ResolveModelWithProviderHint(
|
||||
modelConfig.Model,
|
||||
modelConfig.Provider,
|
||||
)
|
||||
if err != nil {
|
||||
return database.ChatModelConfig{}, "", xerrors.Errorf(
|
||||
"%w: %v",
|
||||
errInvalidModelOverrideMetadata,
|
||||
err,
|
||||
)
|
||||
}
|
||||
return modelConfig, providerName, nil
|
||||
}
|
||||
|
||||
func enabledProviderContainsName(
|
||||
providers []database.AIProvider,
|
||||
providerName string,
|
||||
) bool {
|
||||
normalizedProviderName := chatprovider.NormalizeProvider(providerName)
|
||||
for _, provider := range providers {
|
||||
if chatprovider.NormalizeProvider(string(provider.Type)) == normalizedProviderName {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func userCanUseProviderKeys(
|
||||
providerKeys chatprovider.ProviderAPIKeys,
|
||||
providerName string,
|
||||
@@ -542,18 +509,8 @@ func (p *Server) resolveModelConfigAndNormalizedProvider(
|
||||
}
|
||||
return modelConfig, providerName, nil
|
||||
}
|
||||
modelConfig, providerName, err := validateModelConfigAndResolveProvider(modelConfig)
|
||||
if err != nil {
|
||||
return database.ChatModelConfig{}, "", err
|
||||
}
|
||||
enabledProviders, err := p.configCache.EnabledProviders(ctx)
|
||||
if err != nil {
|
||||
return database.ChatModelConfig{}, "", err
|
||||
}
|
||||
if !enabledProviderContainsName(enabledProviders, providerName) {
|
||||
return database.ChatModelConfig{}, "", sql.ErrNoRows
|
||||
}
|
||||
return modelConfig, providerName, nil
|
||||
// Active configs carry a provider FK; resolved above. Missing FK means no usable config.
|
||||
return database.ChatModelConfig{}, "", sql.ErrNoRows
|
||||
}
|
||||
|
||||
func (p *Server) subagentTools(
|
||||
|
||||
@@ -213,7 +213,6 @@ func seedInternalChatDeps(
|
||||
})
|
||||
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
IsDefault: true,
|
||||
})
|
||||
@@ -580,8 +579,7 @@ func TestResolveChatModel_AIProviderDisabled(t *testing.T) {
|
||||
user, org, _ := seedInternalChatDeps(t, db)
|
||||
provider := insertInternalAIProvider(t, db, database.AIProviderTypeOpenai, "provider-api-key", false)
|
||||
modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o-mini",
|
||||
Model: "gpt-4o-mini",
|
||||
AIProviderID: uuid.NullUUID{
|
||||
UUID: provider.ID,
|
||||
Valid: true,
|
||||
@@ -714,11 +712,30 @@ func insertInternalChatModelConfigWithOptions(
|
||||
) database.ChatModelConfig {
|
||||
t.Helper()
|
||||
|
||||
// Reuse the newest AI provider of this type (creating a bare credential-less
|
||||
// one only when none exists) so the config links the provider already
|
||||
// carrying the test's credentials, or lack thereof, rather than a fresh one.
|
||||
providers, err := db.GetAIProviders(context.Background(), database.GetAIProvidersParams{IncludeDisabled: true})
|
||||
require.NoError(t, err)
|
||||
var aiProvider database.AIProvider
|
||||
for _, candidate := range providers {
|
||||
if candidate.Type != database.AIProviderType(provider) {
|
||||
continue
|
||||
}
|
||||
if aiProvider.ID == uuid.Nil || candidate.CreatedAt.After(aiProvider.CreatedAt) {
|
||||
aiProvider = candidate
|
||||
}
|
||||
}
|
||||
if aiProvider.ID == uuid.Nil {
|
||||
aiProvider = dbgen.AIProvider(t, db, database.AIProvider{
|
||||
Type: database.AIProviderType(provider),
|
||||
})
|
||||
}
|
||||
modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: provider,
|
||||
Model: model,
|
||||
DisplayName: model,
|
||||
Options: options,
|
||||
AIProviderID: uuid.NullUUID{UUID: aiProvider.ID, Valid: true},
|
||||
Model: model,
|
||||
DisplayName: model,
|
||||
Options: options,
|
||||
}, func(p *database.InsertChatModelConfigParams) {
|
||||
p.Enabled = enabled
|
||||
})
|
||||
@@ -1448,7 +1465,6 @@ func TestResolveConfiguredModelOverride_AcceptsAmbientCredentialsProvider(
|
||||
ownerID := uuid.New()
|
||||
modelConfig := database.ChatModelConfig{
|
||||
ID: uuid.New(),
|
||||
Provider: "bedrock",
|
||||
Model: "anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
DisplayName: "Ambient Bedrock Override",
|
||||
Enabled: true,
|
||||
@@ -2115,7 +2131,7 @@ func TestSpawnAgent_ExploreFallsBackWhenOverrideCredentialsAreUnavailable(t *tes
|
||||
currentTurnModel := insertInternalChatModelConfig(
|
||||
t, db, "explore-missing-user-key-current-"+uuid.NewString(), true,
|
||||
)
|
||||
dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
overrideProvider := dbgen.ChatProvider(t, db, database.ChatProvider{
|
||||
Provider: "openai-compat",
|
||||
DisplayName: "OpenAI Compat",
|
||||
}, func(p *database.InsertChatProviderParams) {
|
||||
@@ -2126,9 +2142,9 @@ func TestSpawnAgent_ExploreFallsBackWhenOverrideCredentialsAreUnavailable(t *tes
|
||||
})
|
||||
|
||||
overrideModel := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai-compat",
|
||||
Model: "gpt-4o-mini",
|
||||
DisplayName: "Explore Override Missing User Key",
|
||||
AIProviderID: uuid.NullUUID{UUID: overrideProvider.ID, Valid: true},
|
||||
Model: "gpt-4o-mini",
|
||||
DisplayName: "Explore Override Missing User Key",
|
||||
})
|
||||
require.NoError(t, db.UpsertChatExploreModelOverride(ctx, overrideModel.ID.String()))
|
||||
parentChat := createInternalParentChat(
|
||||
@@ -2764,7 +2780,9 @@ func TestSpawnAgent_ComputerUseUsesComputerUseModelNotParent(t *testing.T) {
|
||||
insertEnabledAnthropicProvider(t, db, user.ID)
|
||||
workspace, build, agent := seedWorkspaceBinding(t, db, user.ID)
|
||||
|
||||
require.Equal(t, "openai", model.Provider, "seed helper must create an OpenAI model")
|
||||
seedProvider, err := db.GetAIProviderByID(ctx, model.AIProviderID.UUID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "openai", string(seedProvider.Type), "seed helper must create an OpenAI model")
|
||||
|
||||
parent, err := server.CreateChat(ctx, CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
@@ -2806,7 +2824,7 @@ func TestSpawnAgent_ComputerUseUsesComputerUseModelNotParent(t *testing.T) {
|
||||
assert.Equal(t, database.ChatModeComputerUse, childChat.Mode.ChatMode)
|
||||
computerUseModelProvider, computerUseModelName, ok := chattool.DefaultComputerUseModel(chattool.ComputerUseProviderAnthropic)
|
||||
require.True(t, ok)
|
||||
assert.NotEqual(t, model.Provider, computerUseModelProvider,
|
||||
assert.NotEqual(t, string(seedProvider.Type), computerUseModelProvider,
|
||||
"computer use model provider must differ from parent model provider")
|
||||
assert.Equal(t, "anthropic", computerUseModelProvider)
|
||||
assert.NotEmpty(t, computerUseModelName)
|
||||
@@ -3668,7 +3686,6 @@ func TestAwaitSubagentCompletion(t *testing.T) {
|
||||
BaseUrl: providerServer.URL,
|
||||
})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Provider: "openai",
|
||||
Model: "gpt-4o-mini",
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
})
|
||||
|
||||
@@ -873,7 +873,7 @@ func newTaskTestFixture(t *testing.T) *taskTestFixture {
|
||||
DisplayName: "openai",
|
||||
BaseUrl: "http://example.invalid",
|
||||
})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{Provider: "openai", IsDefault: true})
|
||||
model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{IsDefault: true})
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
return &taskTestFixture{db: db, pubsub: newTaskRecordingPubsub(ps), sqlDB: sqlDB, user: user, org: org, model: model, apiKey: apiKey}
|
||||
}
|
||||
|
||||
@@ -348,6 +348,8 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback(
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
chat, messages := titleOverrideTestChatAndMessages(t)
|
||||
overrideConfig := titleOverrideModelConfig("gpt-4.1", true)
|
||||
providerID := uuid.New()
|
||||
overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true}
|
||||
|
||||
var requestCount atomic.Int32
|
||||
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
@@ -355,6 +357,12 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback(
|
||||
require.Equal(t, overrideConfig.Model, req.Model)
|
||||
return chattest.OpenAINonStreamingResponse(`{"title":""}`)
|
||||
})
|
||||
provider := database.AIProvider{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
Enabled: true,
|
||||
BaseUrl: serverURL,
|
||||
}
|
||||
keys := titleOverrideOpenAIKeys(serverURL)
|
||||
fallbackModel := &chattest.FakeModel{
|
||||
GenerateObjectFn: func(context.Context, fantasy.ObjectCall) (*fantasy.ObjectResponse, error) {
|
||||
@@ -365,8 +373,11 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback(
|
||||
|
||||
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil)
|
||||
db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil)
|
||||
db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{Type: database.AIProviderTypeOpenai, Enabled: true}}, nil)
|
||||
db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), []uuid.UUID{uuid.Nil}).Return(nil, nil)
|
||||
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
|
||||
ProviderID: providerID,
|
||||
APIKey: "test-key",
|
||||
}}, nil).AnyTimes()
|
||||
|
||||
generated := &generatedChatTitle{}
|
||||
server := titleOverrideTestServer(db, logger)
|
||||
@@ -398,18 +409,34 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnset(t *testing.T) {
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
chat, _ := titleOverrideTestChatAndMessages(t)
|
||||
providerID := uuid.New()
|
||||
preferredConfig := database.ChatModelConfig{
|
||||
ID: uuid.New(),
|
||||
Provider: preferredTitleModels[1].provider,
|
||||
Model: preferredTitleModels[1].model,
|
||||
Enabled: true,
|
||||
ID: uuid.New(),
|
||||
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
|
||||
Model: preferredTitleModels[1].model,
|
||||
Enabled: true,
|
||||
}
|
||||
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
t.Fatal("model construction should not call the provider")
|
||||
return chattest.OpenAIResponse{}
|
||||
})
|
||||
provider := database.AIProvider{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
Enabled: true,
|
||||
BaseUrl: serverURL,
|
||||
}
|
||||
|
||||
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil)
|
||||
db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.ChatModelConfig{
|
||||
{Provider: "openai", Model: "gpt-4.1", Enabled: true},
|
||||
preferredConfig,
|
||||
db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.GetEnabledChatModelConfigsRow{
|
||||
{ChatModelConfig: database.ChatModelConfig{Model: "gpt-4.1", Enabled: true}, Provider: "openai"},
|
||||
{ChatModelConfig: preferredConfig, Provider: preferredTitleModels[1].provider},
|
||||
}, nil)
|
||||
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
|
||||
ProviderID: providerID,
|
||||
APIKey: "test-key",
|
||||
}}, nil).AnyTimes()
|
||||
|
||||
server := titleOverrideTestServer(db, logger)
|
||||
model, gotConfig, _, err := server.resolveManualTitleModel(
|
||||
@@ -435,7 +462,6 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testi
|
||||
providerID := uuid.New()
|
||||
preferredConfig := database.ChatModelConfig{
|
||||
ID: uuid.New(),
|
||||
Provider: preferredTitleModels[1].provider,
|
||||
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
|
||||
Model: preferredTitleModels[1].model,
|
||||
Enabled: true,
|
||||
@@ -452,8 +478,8 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testi
|
||||
}
|
||||
|
||||
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil)
|
||||
db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.ChatModelConfig{
|
||||
preferredConfig,
|
||||
db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.GetEnabledChatModelConfigsRow{
|
||||
{ChatModelConfig: preferredConfig, Provider: preferredTitleModels[1].provider},
|
||||
}, nil)
|
||||
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil)
|
||||
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
|
||||
@@ -483,18 +509,34 @@ func TestResolveManualTitleModel_TitleGenerationOverrideReadDBError(t *testing.T
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
chat, _ := titleOverrideTestChatAndMessages(t)
|
||||
providerID := uuid.New()
|
||||
preferredConfig := database.ChatModelConfig{
|
||||
ID: uuid.New(),
|
||||
Provider: preferredTitleModels[1].provider,
|
||||
Model: preferredTitleModels[1].model,
|
||||
Enabled: true,
|
||||
ID: uuid.New(),
|
||||
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
|
||||
Model: preferredTitleModels[1].model,
|
||||
Enabled: true,
|
||||
}
|
||||
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
t.Fatal("model construction should not call the provider")
|
||||
return chattest.OpenAIResponse{}
|
||||
})
|
||||
provider := database.AIProvider{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
Enabled: true,
|
||||
BaseUrl: serverURL,
|
||||
}
|
||||
|
||||
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", sql.ErrConnDone)
|
||||
db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.ChatModelConfig{
|
||||
{Provider: "openai", Model: "gpt-4.1", Enabled: true},
|
||||
preferredConfig,
|
||||
db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.GetEnabledChatModelConfigsRow{
|
||||
{ChatModelConfig: database.ChatModelConfig{Model: "gpt-4.1", Enabled: true}, Provider: "openai"},
|
||||
{ChatModelConfig: preferredConfig, Provider: preferredTitleModels[1].provider},
|
||||
}, nil)
|
||||
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
|
||||
ProviderID: providerID,
|
||||
APIKey: "test-key",
|
||||
}}, nil).AnyTimes()
|
||||
|
||||
server := titleOverrideTestServer(db, logger)
|
||||
model, gotConfig, _, err := server.resolveManualTitleModel(
|
||||
@@ -518,11 +560,26 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUsable(t *testing.T)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
chat, _ := titleOverrideTestChatAndMessages(t)
|
||||
overrideConfig := titleOverrideModelConfig("gpt-4.1", true)
|
||||
providerID := uuid.New()
|
||||
overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true}
|
||||
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
t.Fatal("model construction should not call the provider")
|
||||
return chattest.OpenAIResponse{}
|
||||
})
|
||||
provider := database.AIProvider{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
Enabled: true,
|
||||
BaseUrl: serverURL,
|
||||
}
|
||||
|
||||
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil)
|
||||
db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil)
|
||||
db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{Type: database.AIProviderTypeOpenai, Enabled: true}}, nil)
|
||||
db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return(nil, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
|
||||
ProviderID: providerID,
|
||||
APIKey: "test-key",
|
||||
}}, nil).AnyTimes()
|
||||
|
||||
server := titleOverrideTestServer(db, logger)
|
||||
model, gotConfig, _, err := server.resolveManualTitleModel(
|
||||
@@ -546,11 +603,18 @@ func TestResolveManualTitleModel_TitleGenerationOverrideMissingCredentials(t *te
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
chat, _ := titleOverrideTestChatAndMessages(t)
|
||||
overrideConfig := titleOverrideModelConfig("gpt-4.1", true)
|
||||
providerID := uuid.New()
|
||||
overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true}
|
||||
provider := database.AIProvider{
|
||||
ID: providerID,
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
Enabled: true,
|
||||
}
|
||||
|
||||
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil)
|
||||
db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil)
|
||||
db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{Type: database.AIProviderTypeOpenai, Enabled: true}}, nil)
|
||||
db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return(nil, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return(nil, nil).AnyTimes()
|
||||
|
||||
server := titleOverrideTestServer(db, logger)
|
||||
model, gotConfig, _, err := server.resolveManualTitleModel(
|
||||
@@ -730,10 +794,9 @@ func titleOverrideTestServer(db database.Store, logger slog.Logger) *Server {
|
||||
|
||||
func titleOverrideModelConfig(model string, enabled bool) database.ChatModelConfig {
|
||||
return database.ChatModelConfig{
|
||||
ID: uuid.New(),
|
||||
Provider: "openai",
|
||||
Model: model,
|
||||
Enabled: enabled,
|
||||
ID: uuid.New(),
|
||||
Model: model,
|
||||
Enabled: enabled,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -43,7 +43,6 @@ func TestUpdateLastTurnSummaryRejectsStaleWrites(t *testing.T) {
|
||||
|
||||
modelCfg, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
DisplayName: "Test Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -138,7 +137,6 @@ func TestPendingChatPersistsSummaryButSkipsWebPush(t *testing.T) {
|
||||
|
||||
modelCfg, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
DisplayName: "Test Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
@@ -222,7 +220,6 @@ func TestSuccessfulChildChatOutcomeSkipsSummaryAndWebPush(t *testing.T) {
|
||||
|
||||
modelCfg, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
Provider: "openai",
|
||||
Model: "test-model",
|
||||
DisplayName: "Test Model",
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
|
||||
Reference in New Issue
Block a user