Thanks to visit codestin.com
Credit goes to github.com

Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions coderd/database/lock.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ const (
LockIDReconcileSystemRoles
LockIDBoundaryUsageStats
LockIDAIProvidersEnvSeed
LockIDChatModelConfigWrites
)

// GenLockID generates a unique and consistent lock ID from a given string.
Expand Down
26 changes: 20 additions & 6 deletions coderd/exp_chats.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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(),
Expand Down
58 changes: 58 additions & 0 deletions coderd/exp_chats_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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()

Expand Down
Loading