mirror of
https://github.com/coder/coder.git
synced 2026-09-22 05:05:20 +08:00
feat(enterprise/dbcrypt): encrypt ai_providers and ai_provider_keys at rest (#25326)
This commit is contained in:
@@ -163,6 +163,52 @@ func Rotate(ctx context.Context, log slog.Logger, sqlDB *sql.DB, ciphers []Ciphe
|
||||
log.Debug(ctx, "encrypted chat provider key", slog.F("provider", provider.Provider), slog.F("current", idx+1), slog.F("cipher", ciphers[0].HexDigest()))
|
||||
}
|
||||
|
||||
aiProviders, err := cryptDB.GetAIProviders(ctx, database.GetAIProvidersParams{IncludeDeleted: true, IncludeDisabled: true})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get ai providers: %w", err)
|
||||
}
|
||||
log.Info(ctx, "encrypting ai provider settings", slog.F("provider_count", len(aiProviders)))
|
||||
for idx, ap := range aiProviders {
|
||||
if !ap.Settings.Valid || strings.TrimSpace(ap.Settings.String) == "" {
|
||||
continue
|
||||
}
|
||||
if ap.SettingsKeyID.Valid && ap.SettingsKeyID.String == ciphers[0].HexDigest() {
|
||||
log.Debug(ctx, "skipping ai provider", slog.F("ai_provider_id", ap.ID), slog.F("name", ap.Name), slog.F("current", idx+1), slog.F("cipher", ciphers[0].HexDigest()))
|
||||
continue
|
||||
}
|
||||
if _, err := cryptDB.UpdateEncryptedAIProviderSettings(ctx, database.UpdateEncryptedAIProviderSettingsParams{
|
||||
ID: ap.ID,
|
||||
Settings: ap.Settings,
|
||||
SettingsKeyID: sql.NullString{}, // dbcrypt will update as required
|
||||
}); err != nil {
|
||||
return xerrors.Errorf("update ai provider id=%s name=%s: %w", ap.ID, ap.Name, err)
|
||||
}
|
||||
log.Debug(ctx, "encrypted ai provider settings", slog.F("ai_provider_id", ap.ID), slog.F("name", ap.Name), slog.F("current", idx+1), slog.F("cipher", ciphers[0].HexDigest()))
|
||||
}
|
||||
|
||||
aiProviderKeys, err := cryptDB.GetAIProviderKeys(ctx)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get ai provider keys: %w", err)
|
||||
}
|
||||
log.Info(ctx, "encrypting ai provider keys", slog.F("key_count", len(aiProviderKeys)))
|
||||
for idx, apk := range aiProviderKeys {
|
||||
if strings.TrimSpace(apk.APIKey) == "" {
|
||||
continue
|
||||
}
|
||||
if apk.ApiKeyKeyID.Valid && apk.ApiKeyKeyID.String == ciphers[0].HexDigest() {
|
||||
log.Debug(ctx, "skipping ai provider key", slog.F("ai_provider_key_id", apk.ID), slog.F("provider_id", apk.ProviderID), slog.F("current", idx+1), slog.F("cipher", ciphers[0].HexDigest()))
|
||||
continue
|
||||
}
|
||||
if _, err := cryptDB.UpdateEncryptedAIProviderKey(ctx, database.UpdateEncryptedAIProviderKeyParams{
|
||||
ID: apk.ID,
|
||||
APIKey: apk.APIKey,
|
||||
ApiKeyKeyID: sql.NullString{}, // dbcrypt will update as required
|
||||
}); err != nil {
|
||||
return xerrors.Errorf("update ai provider key id=%s provider_id=%s: %w", apk.ID, apk.ProviderID, err)
|
||||
}
|
||||
log.Debug(ctx, "encrypted ai provider key", slog.F("ai_provider_key_id", apk.ID), slog.F("provider_id", apk.ProviderID), slog.F("current", idx+1), slog.F("cipher", ciphers[0].HexDigest()))
|
||||
}
|
||||
|
||||
// Revoke old keys
|
||||
for _, c := range ciphers[1:] {
|
||||
if err := db.RevokeDBCryptKey(ctx, c.HexDigest()); err != nil {
|
||||
@@ -326,6 +372,46 @@ func Decrypt(ctx context.Context, log slog.Logger, sqlDB *sql.DB, ciphers []Ciph
|
||||
log.Debug(ctx, "decrypted chat provider key", slog.F("provider", provider.Provider), slog.F("current", idx+1), slog.F("cipher", ciphers[0].HexDigest()))
|
||||
}
|
||||
|
||||
aiProviders, err := cryptDB.GetAIProviders(ctx, database.GetAIProvidersParams{IncludeDeleted: true, IncludeDisabled: true})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get ai providers: %w", err)
|
||||
}
|
||||
log.Info(ctx, "decrypting ai provider settings", slog.F("provider_count", len(aiProviders)))
|
||||
for idx, ap := range aiProviders {
|
||||
if !ap.SettingsKeyID.Valid {
|
||||
log.Debug(ctx, "skipping ai provider", slog.F("ai_provider_id", ap.ID), slog.F("name", ap.Name), slog.F("current", idx+1))
|
||||
continue
|
||||
}
|
||||
if _, err := cryptDB.UpdateEncryptedAIProviderSettings(ctx, database.UpdateEncryptedAIProviderSettingsParams{
|
||||
ID: ap.ID,
|
||||
Settings: ap.Settings,
|
||||
SettingsKeyID: sql.NullString{}, // explicitly clear the key id
|
||||
}); err != nil {
|
||||
return xerrors.Errorf("decrypt ai provider id=%s name=%s: %w", ap.ID, ap.Name, err)
|
||||
}
|
||||
log.Debug(ctx, "decrypted ai provider", slog.F("ai_provider_id", ap.ID), slog.F("name", ap.Name), slog.F("current", idx+1))
|
||||
}
|
||||
|
||||
aiProviderKeys, err := cryptDB.GetAIProviderKeys(ctx)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get ai provider keys: %w", err)
|
||||
}
|
||||
log.Info(ctx, "decrypting ai provider keys", slog.F("key_count", len(aiProviderKeys)))
|
||||
for idx, apk := range aiProviderKeys {
|
||||
if !apk.ApiKeyKeyID.Valid {
|
||||
log.Debug(ctx, "skipping ai provider key", slog.F("ai_provider_key_id", apk.ID), slog.F("provider_id", apk.ProviderID), slog.F("current", idx+1))
|
||||
continue
|
||||
}
|
||||
if _, err := cryptDB.UpdateEncryptedAIProviderKey(ctx, database.UpdateEncryptedAIProviderKeyParams{
|
||||
ID: apk.ID,
|
||||
APIKey: apk.APIKey,
|
||||
ApiKeyKeyID: sql.NullString{}, // explicitly clear the key id
|
||||
}); err != nil {
|
||||
return xerrors.Errorf("decrypt ai provider key id=%s provider_id=%s: %w", apk.ID, apk.ProviderID, err)
|
||||
}
|
||||
log.Debug(ctx, "decrypted ai provider key", slog.F("ai_provider_key_id", apk.ID), slog.F("provider_id", apk.ProviderID), slog.F("current", idx+1))
|
||||
}
|
||||
|
||||
// Revoke _all_ keys
|
||||
for _, c := range ciphers {
|
||||
if err := db.RevokeDBCryptKey(ctx, c.HexDigest()); err != nil {
|
||||
@@ -354,6 +440,12 @@ UPDATE chat_providers
|
||||
SET api_key = '',
|
||||
api_key_key_id = NULL
|
||||
WHERE api_key_key_id IS NOT NULL;
|
||||
UPDATE ai_providers
|
||||
SET settings = NULL,
|
||||
settings_key_id = NULL
|
||||
WHERE settings_key_id IS NOT NULL;
|
||||
DELETE FROM ai_provider_keys
|
||||
WHERE api_key_key_id IS NOT NULL;
|
||||
COMMIT;
|
||||
`
|
||||
|
||||
|
||||
@@ -385,6 +385,197 @@ func (db *dbCrypt) GetCryptoKeysByFeature(ctx context.Context, feature database.
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
// decryptAIProvider decrypts the secret fields of an AI provider row.
|
||||
func (db *dbCrypt) decryptAIProvider(p *database.AIProvider) error {
|
||||
if !p.Settings.Valid {
|
||||
return nil
|
||||
}
|
||||
return db.decryptField(&p.Settings.String, p.SettingsKeyID)
|
||||
}
|
||||
|
||||
// decryptAIProviderKey decrypts the api_key field of an AI provider key row.
|
||||
func (db *dbCrypt) decryptAIProviderKey(k *database.AIProviderKey) error {
|
||||
return db.decryptField(&k.APIKey, k.ApiKeyKeyID)
|
||||
}
|
||||
|
||||
// encryptAIProviderSettings encrypts the settings column in place,
|
||||
// updating settings_key_id as a side effect. A NULL or blank settings
|
||||
// value clears any associated key reference.
|
||||
func (db *dbCrypt) encryptAIProviderSettings(settings *sql.NullString, keyID *sql.NullString) error {
|
||||
if !settings.Valid || strings.TrimSpace(settings.String) == "" {
|
||||
*settings = sql.NullString{}
|
||||
*keyID = sql.NullString{}
|
||||
return nil
|
||||
}
|
||||
return db.encryptField(&settings.String, keyID)
|
||||
}
|
||||
|
||||
func (db *dbCrypt) GetAIProviderByID(ctx context.Context, id uuid.UUID) (database.AIProvider, error) {
|
||||
provider, err := db.Store.GetAIProviderByID(ctx, id)
|
||||
if err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
if err := db.decryptAIProvider(&provider); err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
return provider, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) GetAIProviderByName(ctx context.Context, name string) (database.AIProvider, error) {
|
||||
provider, err := db.Store.GetAIProviderByName(ctx, name)
|
||||
if err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
if err := db.decryptAIProvider(&provider); err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
return provider, nil
|
||||
}
|
||||
|
||||
// GetAIProviders returns AI provider rows, with their settings
|
||||
// decrypted, honoring the include_deleted and include_disabled flags
|
||||
// from the underlying query.
|
||||
func (db *dbCrypt) GetAIProviders(ctx context.Context, arg database.GetAIProvidersParams) ([]database.AIProvider, error) {
|
||||
providers, err := db.Store.GetAIProviders(ctx, arg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i := range providers {
|
||||
if err := db.decryptAIProvider(&providers[i]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return providers, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) InsertAIProvider(ctx context.Context, params database.InsertAIProviderParams) (database.AIProvider, error) {
|
||||
if err := db.encryptAIProviderSettings(¶ms.Settings, ¶ms.SettingsKeyID); err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
|
||||
provider, err := db.Store.InsertAIProvider(ctx, params)
|
||||
if err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
if err := db.decryptAIProvider(&provider); err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
return provider, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) UpdateAIProvider(ctx context.Context, params database.UpdateAIProviderParams) (database.AIProvider, error) {
|
||||
if err := db.encryptAIProviderSettings(¶ms.Settings, ¶ms.SettingsKeyID); err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
|
||||
provider, err := db.Store.UpdateAIProvider(ctx, params)
|
||||
if err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
if err := db.decryptAIProvider(&provider); err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
return provider, nil
|
||||
}
|
||||
|
||||
// UpdateEncryptedAIProviderSettings re-encrypts the settings column
|
||||
// of a row, regardless of its deleted flag, so that dbcrypt key
|
||||
// rotation can move every FK reference to a new key digest before
|
||||
// old keys are revoked.
|
||||
func (db *dbCrypt) UpdateEncryptedAIProviderSettings(ctx context.Context, params database.UpdateEncryptedAIProviderSettingsParams) (database.AIProvider, error) {
|
||||
if err := db.encryptAIProviderSettings(¶ms.Settings, ¶ms.SettingsKeyID); err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
|
||||
provider, err := db.Store.UpdateEncryptedAIProviderSettings(ctx, params)
|
||||
if err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
if err := db.decryptAIProvider(&provider); err != nil {
|
||||
return database.AIProvider{}, err
|
||||
}
|
||||
return provider, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) GetAIProviderKeyByID(ctx context.Context, id uuid.UUID) (database.AIProviderKey, error) {
|
||||
key, err := db.Store.GetAIProviderKeyByID(ctx, id)
|
||||
if err != nil {
|
||||
return database.AIProviderKey{}, err
|
||||
}
|
||||
if err := db.decryptAIProviderKey(&key); err != nil {
|
||||
return database.AIProviderKey{}, err
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) GetAIProviderKeysByProviderID(ctx context.Context, providerID uuid.UUID) ([]database.AIProviderKey, error) {
|
||||
keys, err := db.Store.GetAIProviderKeysByProviderID(ctx, providerID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i := range keys {
|
||||
if err := db.decryptAIProviderKey(&keys[i]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) InsertAIProviderKey(ctx context.Context, params database.InsertAIProviderKeyParams) (database.AIProviderKey, error) {
|
||||
if strings.TrimSpace(params.APIKey) == "" {
|
||||
params.ApiKeyKeyID = sql.NullString{}
|
||||
} else if err := db.encryptField(¶ms.APIKey, ¶ms.ApiKeyKeyID); err != nil {
|
||||
return database.AIProviderKey{}, err
|
||||
}
|
||||
|
||||
key, err := db.Store.InsertAIProviderKey(ctx, params)
|
||||
if err != nil {
|
||||
return database.AIProviderKey{}, err
|
||||
}
|
||||
if err := db.decryptAIProviderKey(&key); err != nil {
|
||||
return database.AIProviderKey{}, err
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
// GetAIProviderKeys returns every AI provider key row, including
|
||||
// those whose provider has been soft-deleted, with their api_key
|
||||
// decrypted. The dbcrypt key rotation utility uses this to walk
|
||||
// every row holding a foreign-key reference to dbcrypt_keys before
|
||||
// old keys are revoked.
|
||||
func (db *dbCrypt) GetAIProviderKeys(ctx context.Context) ([]database.AIProviderKey, error) {
|
||||
keys, err := db.Store.GetAIProviderKeys(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i := range keys {
|
||||
if err := db.decryptAIProviderKey(&keys[i]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
// UpdateEncryptedAIProviderKey re-encrypts the api_key column of a
|
||||
// key row, so that dbcrypt key rotation can move every FK reference
|
||||
// to a new key digest before old keys are revoked.
|
||||
func (db *dbCrypt) UpdateEncryptedAIProviderKey(ctx context.Context, params database.UpdateEncryptedAIProviderKeyParams) (database.AIProviderKey, error) {
|
||||
if strings.TrimSpace(params.APIKey) == "" {
|
||||
params.ApiKeyKeyID = sql.NullString{}
|
||||
} else if err := db.encryptField(¶ms.APIKey, ¶ms.ApiKeyKeyID); err != nil {
|
||||
return database.AIProviderKey{}, err
|
||||
}
|
||||
|
||||
key, err := db.Store.UpdateEncryptedAIProviderKey(ctx, params)
|
||||
if err != nil {
|
||||
return database.AIProviderKey{}, err
|
||||
}
|
||||
if err := db.decryptAIProviderKey(&key); err != nil {
|
||||
return database.AIProviderKey{}, err
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) GetChatProviderByID(ctx context.Context, id uuid.UUID) (database.ChatProvider, error) {
|
||||
provider, err := db.Store.GetChatProviderByID(ctx, id)
|
||||
if err != nil {
|
||||
|
||||
@@ -1055,6 +1055,243 @@ func TestMCPServerConfigs(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func requireAIProviderDecrypted(
|
||||
t *testing.T,
|
||||
provider database.AIProvider,
|
||||
ciphers []Cipher,
|
||||
wantSettings string,
|
||||
) {
|
||||
t.Helper()
|
||||
if wantSettings == "" {
|
||||
require.False(t, provider.Settings.Valid)
|
||||
require.False(t, provider.SettingsKeyID.Valid)
|
||||
return
|
||||
}
|
||||
require.True(t, provider.Settings.Valid)
|
||||
require.Equal(t, wantSettings, provider.Settings.String)
|
||||
require.Equal(t, ciphers[0].HexDigest(), provider.SettingsKeyID.String)
|
||||
}
|
||||
|
||||
func requireAIProviderRawEncrypted(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
rawDB database.Store,
|
||||
providerID uuid.UUID,
|
||||
ciphers []Cipher,
|
||||
wantSettings string,
|
||||
) {
|
||||
t.Helper()
|
||||
raw, err := rawDB.GetAIProviderByID(ctx, providerID)
|
||||
require.NoError(t, err)
|
||||
require.True(t, raw.Settings.Valid)
|
||||
requireEncryptedEquals(t, ciphers[0], raw.Settings.String, wantSettings)
|
||||
}
|
||||
|
||||
func requireAIProviderKeyDecrypted(
|
||||
t *testing.T,
|
||||
key database.AIProviderKey,
|
||||
ciphers []Cipher,
|
||||
wantAPIKey string,
|
||||
) {
|
||||
t.Helper()
|
||||
require.Equal(t, wantAPIKey, key.APIKey)
|
||||
if wantAPIKey != "" {
|
||||
require.Equal(t, ciphers[0].HexDigest(), key.ApiKeyKeyID.String)
|
||||
} else {
|
||||
require.False(t, key.ApiKeyKeyID.Valid)
|
||||
}
|
||||
}
|
||||
|
||||
func requireAIProviderKeyRawEncrypted(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
rawDB database.Store,
|
||||
keyID uuid.UUID,
|
||||
ciphers []Cipher,
|
||||
wantAPIKey string,
|
||||
) {
|
||||
t.Helper()
|
||||
raw, err := rawDB.GetAIProviderKeyByID(ctx, keyID)
|
||||
require.NoError(t, err)
|
||||
requireEncryptedEquals(t, ciphers[0], raw.APIKey, wantAPIKey)
|
||||
}
|
||||
|
||||
func TestAIProviders(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
|
||||
//nolint:gosec // test fixture, not real credentials
|
||||
const settings = `{"_type":"bedrock","_version":1,"region":"us-west-2","model":"anthropic.claude-sonnet-4-5-20250929-v1:0","access_key":"AKIA-test","access_key_secret":"test-secret"}`
|
||||
|
||||
insertProvider := func(t *testing.T, crypt *dbCrypt, ciphers []Cipher) database.AIProvider {
|
||||
t.Helper()
|
||||
provider := dbgen.AIProvider(t, crypt, database.AIProvider{
|
||||
Name: "anthropic-bedrock",
|
||||
Type: database.AiProviderTypeAnthropic,
|
||||
BaseUrl: "https://bedrock-runtime.us-west-2.amazonaws.com/",
|
||||
Settings: sql.NullString{String: settings, Valid: true},
|
||||
})
|
||||
requireAIProviderDecrypted(t, provider, ciphers, settings)
|
||||
return provider
|
||||
}
|
||||
|
||||
t.Run("InsertAIProvider", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider := insertProvider(t, crypt, ciphers)
|
||||
requireAIProviderRawEncrypted(ctx, t, db, provider.ID, ciphers, settings)
|
||||
})
|
||||
|
||||
t.Run("InsertAIProviderEmptySettings", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, _ := setup(t)
|
||||
provider := dbgen.AIProvider(t, crypt, database.AIProvider{
|
||||
Name: "openai-empty",
|
||||
}, func(p *database.InsertAIProviderParams) {
|
||||
p.Settings = sql.NullString{}
|
||||
})
|
||||
require.False(t, provider.SettingsKeyID.Valid)
|
||||
raw, err := db.GetAIProviderByID(ctx, provider.ID)
|
||||
require.NoError(t, err)
|
||||
require.False(t, raw.Settings.Valid)
|
||||
})
|
||||
|
||||
t.Run("GetAIProviderByID", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider := insertProvider(t, crypt, ciphers)
|
||||
got, err := crypt.GetAIProviderByID(ctx, provider.ID)
|
||||
require.NoError(t, err)
|
||||
requireAIProviderDecrypted(t, got, ciphers, settings)
|
||||
requireAIProviderRawEncrypted(ctx, t, db, provider.ID, ciphers, settings)
|
||||
})
|
||||
|
||||
t.Run("GetAIProviderByName", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider := insertProvider(t, crypt, ciphers)
|
||||
got, err := crypt.GetAIProviderByName(ctx, provider.Name)
|
||||
require.NoError(t, err)
|
||||
requireAIProviderDecrypted(t, got, ciphers, settings)
|
||||
requireAIProviderRawEncrypted(ctx, t, db, provider.ID, ciphers, settings)
|
||||
})
|
||||
|
||||
t.Run("GetAIProviders", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider := insertProvider(t, crypt, ciphers)
|
||||
providers, err := crypt.GetAIProviders(ctx, database.GetAIProvidersParams{})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, providers, 1)
|
||||
requireAIProviderDecrypted(t, providers[0], ciphers, settings)
|
||||
requireAIProviderRawEncrypted(ctx, t, db, provider.ID, ciphers, settings)
|
||||
})
|
||||
|
||||
t.Run("UpdateAIProvider", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider := insertProvider(t, crypt, ciphers)
|
||||
//nolint:gosec // test fixture, not real credentials
|
||||
const newSettings = `{"_type":"bedrock","_version":1,"region":"us-east-1","model":"anthropic.claude-sonnet-4-5-20250929-v1:0","access_key":"AKIA-test","access_key_secret":"test-secret"}`
|
||||
updated, err := crypt.UpdateAIProvider(ctx, database.UpdateAIProviderParams{
|
||||
ID: provider.ID,
|
||||
DisplayName: provider.DisplayName,
|
||||
Enabled: provider.Enabled,
|
||||
BaseUrl: provider.BaseUrl,
|
||||
Settings: sql.NullString{String: newSettings, Valid: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
requireAIProviderDecrypted(t, updated, ciphers, newSettings)
|
||||
requireAIProviderRawEncrypted(ctx, t, db, provider.ID, ciphers, newSettings)
|
||||
})
|
||||
|
||||
t.Run("UpdateAIProviderClearsSettings", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider := insertProvider(t, crypt, ciphers)
|
||||
updated, err := crypt.UpdateAIProvider(ctx, database.UpdateAIProviderParams{
|
||||
ID: provider.ID,
|
||||
DisplayName: provider.DisplayName,
|
||||
Enabled: provider.Enabled,
|
||||
BaseUrl: provider.BaseUrl,
|
||||
Settings: sql.NullString{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, updated.SettingsKeyID.Valid)
|
||||
raw, err := db.GetAIProviderByID(ctx, provider.ID)
|
||||
require.NoError(t, err)
|
||||
require.False(t, raw.Settings.Valid)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAIProviderKeys(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
|
||||
//nolint:gosec // test credentials
|
||||
const apiKey = "sk-test-api-key"
|
||||
|
||||
insertProviderAndKey := func(t *testing.T, crypt *dbCrypt, ciphers []Cipher) (database.AIProvider, database.AIProviderKey) {
|
||||
t.Helper()
|
||||
provider := dbgen.AIProvider(t, crypt, database.AIProvider{
|
||||
Name: "openai-test",
|
||||
Type: database.AiProviderTypeOpenai,
|
||||
BaseUrl: "https://api.openai.com/v1/",
|
||||
})
|
||||
key := dbgen.AIProviderKey(t, crypt, database.AIProviderKey{
|
||||
ProviderID: provider.ID,
|
||||
APIKey: apiKey,
|
||||
})
|
||||
requireAIProviderKeyDecrypted(t, key, ciphers, apiKey)
|
||||
return provider, key
|
||||
}
|
||||
|
||||
t.Run("InsertAIProviderKey", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
_, key := insertProviderAndKey(t, crypt, ciphers)
|
||||
requireAIProviderKeyRawEncrypted(ctx, t, db, key.ID, ciphers, apiKey)
|
||||
})
|
||||
|
||||
t.Run("InsertAIProviderKeyEmpty", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, _ := setup(t)
|
||||
provider := dbgen.AIProvider(t, crypt, database.AIProvider{
|
||||
Name: "openai-empty-key",
|
||||
})
|
||||
key := dbgen.AIProviderKey(t, crypt, database.AIProviderKey{
|
||||
ProviderID: provider.ID,
|
||||
}, func(p *database.InsertAIProviderKeyParams) {
|
||||
p.APIKey = ""
|
||||
})
|
||||
require.False(t, key.ApiKeyKeyID.Valid)
|
||||
raw, err := db.GetAIProviderKeyByID(ctx, key.ID)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, raw.APIKey)
|
||||
})
|
||||
|
||||
t.Run("GetAIProviderKeysByProviderID", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider, key := insertProviderAndKey(t, crypt, ciphers)
|
||||
keys, err := crypt.GetAIProviderKeysByProviderID(ctx, provider.ID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, keys, 1)
|
||||
requireAIProviderKeyDecrypted(t, keys[0], ciphers, apiKey)
|
||||
requireAIProviderKeyRawEncrypted(ctx, t, db, key.ID, ciphers, apiKey)
|
||||
})
|
||||
|
||||
t.Run("DeleteAIProviderKey", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider, key := insertProviderAndKey(t, crypt, ciphers)
|
||||
require.NoError(t, crypt.DeleteAIProviderKey(ctx, key.ID))
|
||||
keys, err := db.GetAIProviderKeysByProviderID(ctx, provider.ID)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, keys)
|
||||
})
|
||||
}
|
||||
|
||||
func TestMCPServerUserTokens(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
|
||||
Reference in New Issue
Block a user