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:
Mathias Fredriksson
2026-07-01 15:59:55 +03:00
committed by GitHub
parent cf75dd0f46
commit 047c47495b
92 changed files with 1072 additions and 1197 deletions
-26
View File
@@ -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))
}
}
}
-120
View File
@@ -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))
})
}
-1
View File
@@ -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,
+1 -15
View File
@@ -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
}
+5 -20
View File
@@ -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{
+4 -4
View File
@@ -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,
+10 -6
View File
@@ -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
View File
@@ -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())
+2 -32
View File
@@ -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
}
-2
View File
@@ -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,
})
-5
View File
@@ -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;
@@ -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'
);
-1
View File
@@ -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"`
+2 -8
View File
@@ -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
+27 -52
View File
@@ -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},
+65 -114
View File
@@ -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 {
+12 -44
View File
@@ -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
+9 -5
View File
@@ -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
View File
@@ -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
View File
@@ -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(
+9 -9
View File
@@ -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,
+22 -12
View File
@@ -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
View File
@@ -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),
)
-2
View File
@@ -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,
+93 -8
View File
@@ -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)
}
-1
View File
@@ -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},
+32 -17
View File
@@ -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,
+23 -19
View File
@@ -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)
+2 -2
View File
@@ -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,
-1
View File
@@ -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
-1
View File
@@ -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})
-1
View File
@@ -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,
+5 -1
View File
@@ -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},
-1
View File
@@ -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})
-3
View File
@@ -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 -1
View File
@@ -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},
}
+20 -5
View File
@@ -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,
+8 -10
View File
@@ -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)
+2 -45
View File
@@ -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(
+32 -15
View File
@@ -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},
})
+1 -1
View File
@@ -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}
}
+90 -27
View File
@@ -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},