diff --git a/enterprise/dbcrypt/cliutil.go b/enterprise/dbcrypt/cliutil.go index 84a2a2344a..ef85fde2cb 100644 --- a/enterprise/dbcrypt/cliutil.go +++ b/enterprise/dbcrypt/cliutil.go @@ -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; ` diff --git a/enterprise/dbcrypt/dbcrypt.go b/enterprise/dbcrypt/dbcrypt.go index a222de1607..bc0e231aa1 100644 --- a/enterprise/dbcrypt/dbcrypt.go +++ b/enterprise/dbcrypt/dbcrypt.go @@ -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 { diff --git a/enterprise/dbcrypt/dbcrypt_internal_test.go b/enterprise/dbcrypt/dbcrypt_internal_test.go index f6d24270d7..fea3a4eeb6 100644 --- a/enterprise/dbcrypt/dbcrypt_internal_test.go +++ b/enterprise/dbcrypt/dbcrypt_internal_test.go @@ -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()