diff --git a/coderd/database/lock.go b/coderd/database/lock.go index 8d0894abc87..d2ec69293dc 100644 --- a/coderd/database/lock.go +++ b/coderd/database/lock.go @@ -16,6 +16,7 @@ const ( LockIDReconcileSystemRoles LockIDBoundaryUsageStats LockIDAIProvidersEnvSeed + LockIDChatModelConfigWrites ) // GenLockID generates a unique and consistent lock ID from a given string. diff --git a/coderd/exp_chats.go b/coderd/exp_chats.go index fa4407f880e..e5f350fa34c 100644 --- a/coderd/exp_chats.go +++ b/coderd/exp_chats.go @@ -7039,6 +7039,20 @@ func validateChatModelConfigProviderModel(aiProvider database.AIProvider, model return nil } +// inChatModelConfigWriteTx runs fn in a transaction that holds the advisory +// lock serializing chat model config writes. All writes to the table must go +// through this helper so concurrent writers cannot act on stale reads, e.g. +// two creates on an empty deployment both self-promoting to default and +// violating the idx_chat_model_configs_single_default unique index. +func (api *API) inChatModelConfigWriteTx(ctx context.Context, fn func(tx database.Store) error) error { + return api.Database.InTx(func(tx database.Store) error { + if err := tx.AcquireLock(ctx, database.LockIDChatModelConfigWrites); err != nil { + return xerrors.Errorf("acquire chat model config write lock: %w", err) + } + return fn(tx) + }, nil) +} + func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) { ctx := r.Context() apiKey := httpmw.APIKey(r) @@ -7141,7 +7155,7 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) { } var inserted database.ChatModelConfig - err = api.Database.InTx(func(tx database.Store) error { + err = api.inChatModelConfigWriteTx(ctx, func(tx database.Store) error { //nolint:gocritic // The route already authorized chat model config updates. lockedAIProvider, err := tx.GetAIProviderByIDForReferenceLock(dbauthz.AsChatd(ctx), insertParams.AIProviderID.UUID) if err != nil { @@ -7193,7 +7207,7 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) { } inserted = refreshedConfig return nil - }, nil) + }) if err != nil { var providerModelErr *chatModelConfigProviderModelError switch { @@ -7350,7 +7364,7 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) { // Re-derive the provider type under lock when the model or provider changes. revalidateProviderModel := updateParams.AIProviderID.Valid && (req.AIProviderID != nil || strings.TrimSpace(req.Model) != "") var updated database.ChatModelConfig - err = api.Database.InTx(func(tx database.Store) error { + err = api.inChatModelConfigWriteTx(ctx, func(tx database.Store) error { if revalidateProviderModel { //nolint:gocritic // The route already authorized chat model config updates. aiProvider, err := tx.GetAIProviderByIDForReferenceLock(dbauthz.AsChatd(ctx), updateParams.AIProviderID.UUID) @@ -7406,7 +7420,7 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) { } updated = refreshedConfig return nil - }, nil) + }) if err != nil { var providerModelErr *chatModelConfigProviderModelError switch { @@ -7466,12 +7480,12 @@ func (api *API) deleteChatModelConfig(rw http.ResponseWriter, r *http.Request) { return } - if err := api.Database.InTx(func(tx database.Store) error { + if err := api.inChatModelConfigWriteTx(ctx, func(tx database.Store) error { if err := tx.DeleteChatModelConfigByID(ctx, modelConfigID); err != nil { return err } return ensureDefaultChatModelConfig(ctx, tx) - }, nil); err != nil { + }); err != nil { httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ Message: "Failed to delete chat model config.", Detail: err.Error(), diff --git a/coderd/exp_chats_test.go b/coderd/exp_chats_test.go index 4b93cfa1b8a..f40c155e476 100644 --- a/coderd/exp_chats_test.go +++ b/coderd/exp_chats_test.go @@ -23,6 +23,7 @@ import ( "github.com/shopspring/decimal" "github.com/sqlc-dev/pqtype" "github.com/stretchr/testify/require" + "golang.org/x/sync/errgroup" "golang.org/x/xerrors" "cdr.dev/slog/v3/sloggers/slogtest" @@ -3789,6 +3790,63 @@ func TestCreateChatModelConfig(t *testing.T) { requireChatModelPricing(t, configs[0].ModelConfig, pricing) }) + t.Run("ConcurrentCreatesElectSingleDefault", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client := newChatClient(t) + _ = coderdtest.CreateFirstUser(t, client.Client) + + aiProvider := createAIProviderForTest(t, client, "openai", "test-api-key") + + // Concurrent creators race to self-elect a default while one claims + // it via a follow-up update, mirroring a terraform apply. Unserialized, + // the losers 409 on the single-default unique index. + const creators = 10 + contextLimit := int64(4096) + var claimed codersdk.ChatModelConfig + var eg errgroup.Group + for i := range creators - 1 { + eg.Go(func() error { + _, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ + AIProviderID: &aiProvider.ID, + Model: fmt.Sprintf("gpt-4o-mini-%d", i), + ContextLimit: &contextLimit, + }) + return err + }) + } + eg.Go(func() error { + created, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ + AIProviderID: &aiProvider.ID, + Model: "gpt-4o", + ContextLimit: &contextLimit, + }) + if err != nil { + return xerrors.Errorf("create claimed config: %w", err) + } + claimed, err = client.UpdateChatModelConfig(ctx, created.ID, codersdk.UpdateChatModelConfigRequest{ + IsDefault: ptr.Ref(true), + }) + if err != nil { + return xerrors.Errorf("promote claimed config: %w", err) + } + return nil + }) + require.NoError(t, eg.Wait()) + + configs, err := client.ListChatModelConfigs(ctx) + require.NoError(t, err) + require.Len(t, configs, creators) + var defaults []uuid.UUID + for _, cfg := range configs { + if cfg.IsDefault { + defaults = append(defaults, cfg.ID) + } + } + require.Equal(t, []uuid.UUID{claimed.ID}, defaults) + }) + t.Run("RejectsNegativePricing", func(t *testing.T) { t.Parallel()