mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add AI provider schema expansion (#25412)
This commit is contained in:
@@ -5,6 +5,7 @@ import (
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/lib/pq"
|
||||
@@ -234,6 +235,25 @@ func genData(t *testing.T, db database.Store) []database.User {
|
||||
OAuthAccessToken: "access-" + usr.ID.String(),
|
||||
OAuthRefreshToken: "refresh-" + usr.ID.String(),
|
||||
})
|
||||
provider := dbgen.AIProvider(t, db, database.AIProvider{
|
||||
Name: "ai-provider-" + usr.ID.String(),
|
||||
Settings: sql.NullString{String: "settings-" + usr.ID.String(), Valid: true},
|
||||
})
|
||||
_ = dbgen.AIProviderKey(t, db, database.AIProviderKey{
|
||||
ProviderID: provider.ID,
|
||||
APIKey: "provider-key-" + usr.ID.String(),
|
||||
})
|
||||
now := time.Now()
|
||||
_, err := db.UpsertUserAIProviderKey(context.Background(), database.UpsertUserAIProviderKeyParams{
|
||||
ID: uuid.New(),
|
||||
UserID: usr.ID,
|
||||
AIProviderID: provider.ID,
|
||||
APIKey: "user-ai-provider-key-" + usr.ID.String(),
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Deleted users cannot have user_links or user_secrets.
|
||||
if !deleted {
|
||||
// Fun fact: our schema allows _all_ login types to have
|
||||
@@ -302,6 +322,36 @@ func requireEncryptedWithCipher(ctx context.Context, t *testing.T, db database.S
|
||||
requireEncryptedEquals(t, c, "value-"+userID.String(), s.Value)
|
||||
require.Equal(t, c.HexDigest(), s.ValueKeyID.String)
|
||||
}
|
||||
|
||||
providers, err := db.GetAIProviders(ctx, database.GetAIProvidersParams{
|
||||
IncludeDeleted: true,
|
||||
IncludeDisabled: true,
|
||||
})
|
||||
require.NoError(t, err, "failed to get ai providers")
|
||||
providerName := "ai-provider-" + userID.String()
|
||||
var provider database.AIProvider
|
||||
for _, p := range providers {
|
||||
if p.Name == providerName {
|
||||
provider = p
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotEqual(t, uuid.Nil, provider.ID, "expected ai provider for user %s", userID)
|
||||
require.True(t, provider.Settings.Valid)
|
||||
requireEncryptedEquals(t, c, "settings-"+userID.String(), provider.Settings.String)
|
||||
require.Equal(t, c.HexDigest(), provider.SettingsKeyID.String)
|
||||
|
||||
providerKeys, err := db.GetAIProviderKeysByProviderID(ctx, provider.ID)
|
||||
require.NoError(t, err, "failed to get ai provider keys for provider %s", provider.ID)
|
||||
require.Len(t, providerKeys, 1)
|
||||
requireEncryptedEquals(t, c, "provider-key-"+userID.String(), providerKeys[0].APIKey)
|
||||
require.Equal(t, c.HexDigest(), providerKeys[0].ApiKeyKeyID.String)
|
||||
|
||||
userAIProviderKeys, err := db.GetUserAIProviderKeysByUserID(ctx, userID)
|
||||
require.NoError(t, err, "failed to get user ai provider keys for user %s", userID)
|
||||
require.Len(t, userAIProviderKeys, 1)
|
||||
requireEncryptedEquals(t, c, "user-ai-provider-key-"+userID.String(), userAIProviderKeys[0].APIKey)
|
||||
require.Equal(t, c.HexDigest(), userAIProviderKeys[0].ApiKeyKeyID.String)
|
||||
}
|
||||
|
||||
// nullCipher is a dbcrypt.Cipher that does not encrypt or decrypt.
|
||||
|
||||
@@ -209,6 +209,29 @@ func Rotate(ctx context.Context, log slog.Logger, sqlDB *sql.DB, ciphers []Ciphe
|
||||
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()))
|
||||
}
|
||||
|
||||
userAIProviderKeys, err := cryptDB.GetUserAIProviderKeys(ctx)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get user ai provider keys: %w", err)
|
||||
}
|
||||
log.Info(ctx, "encrypting user ai provider keys", slog.F("key_count", len(userAIProviderKeys)))
|
||||
for idx, key := range userAIProviderKeys {
|
||||
if strings.TrimSpace(key.APIKey) == "" {
|
||||
continue
|
||||
}
|
||||
if key.ApiKeyKeyID.Valid && key.ApiKeyKeyID.String == ciphers[0].HexDigest() {
|
||||
log.Debug(ctx, "skipping user ai provider key", slog.F("user_ai_provider_key_id", key.ID), slog.F("ai_provider_id", key.AIProviderID), slog.F("user_id", key.UserID), slog.F("current", idx+1), slog.F("cipher", ciphers[0].HexDigest()))
|
||||
continue
|
||||
}
|
||||
if _, err := cryptDB.UpdateEncryptedUserAIProviderKey(ctx, database.UpdateEncryptedUserAIProviderKeyParams{
|
||||
ID: key.ID,
|
||||
APIKey: key.APIKey,
|
||||
ApiKeyKeyID: sql.NullString{}, // dbcrypt will update as required
|
||||
}); err != nil {
|
||||
return xerrors.Errorf("update user ai provider key id=%s ai_provider_id=%s user_id=%s: %w", key.ID, key.AIProviderID, key.UserID, err)
|
||||
}
|
||||
log.Debug(ctx, "encrypted user ai provider key", slog.F("user_ai_provider_key_id", key.ID), slog.F("ai_provider_id", key.AIProviderID), slog.F("user_id", key.UserID), 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 {
|
||||
@@ -412,6 +435,26 @@ func Decrypt(ctx context.Context, log slog.Logger, sqlDB *sql.DB, ciphers []Ciph
|
||||
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))
|
||||
}
|
||||
|
||||
userAIProviderKeys, err := cryptDB.GetUserAIProviderKeys(ctx)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get user ai provider keys: %w", err)
|
||||
}
|
||||
log.Info(ctx, "decrypting user ai provider keys", slog.F("key_count", len(userAIProviderKeys)))
|
||||
for idx, key := range userAIProviderKeys {
|
||||
if !key.ApiKeyKeyID.Valid {
|
||||
log.Debug(ctx, "skipping user ai provider key", slog.F("user_ai_provider_key_id", key.ID), slog.F("ai_provider_id", key.AIProviderID), slog.F("user_id", key.UserID), slog.F("current", idx+1))
|
||||
continue
|
||||
}
|
||||
if _, err := cryptDB.UpdateEncryptedUserAIProviderKey(ctx, database.UpdateEncryptedUserAIProviderKeyParams{
|
||||
ID: key.ID,
|
||||
APIKey: key.APIKey,
|
||||
ApiKeyKeyID: sql.NullString{}, // explicitly clear the key id
|
||||
}); err != nil {
|
||||
return xerrors.Errorf("decrypt user ai provider key id=%s ai_provider_id=%s user_id=%s: %w", key.ID, key.AIProviderID, key.UserID, err)
|
||||
}
|
||||
log.Debug(ctx, "decrypted user ai provider key", slog.F("user_ai_provider_key_id", key.ID), slog.F("ai_provider_id", key.AIProviderID), slog.F("user_id", key.UserID), slog.F("current", idx+1))
|
||||
}
|
||||
|
||||
// Revoke _all_ keys
|
||||
for _, c := range ciphers {
|
||||
if err := db.RevokeDBCryptKey(ctx, c.HexDigest()); err != nil {
|
||||
@@ -434,6 +477,8 @@ DELETE FROM external_auth_links
|
||||
OR oauth_refresh_token_key_id IS NOT NULL;
|
||||
DELETE FROM user_chat_provider_keys
|
||||
WHERE api_key_key_id IS NOT NULL;
|
||||
DELETE FROM user_ai_provider_keys
|
||||
WHERE api_key_key_id IS NOT NULL;
|
||||
DELETE FROM user_secrets
|
||||
WHERE value_key_id IS NOT NULL;
|
||||
UPDATE chat_providers
|
||||
|
||||
@@ -662,6 +662,98 @@ func (db *dbCrypt) UpdateChatProvider(ctx context.Context, params database.Updat
|
||||
return provider, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) decryptUserAIProviderKey(key *database.UserAiProviderKey) error {
|
||||
return db.decryptField(&key.APIKey, key.ApiKeyKeyID)
|
||||
}
|
||||
|
||||
func (db *dbCrypt) GetUserAIProviderKeyByProviderID(ctx context.Context, params database.GetUserAIProviderKeyByProviderIDParams) (database.UserAiProviderKey, error) {
|
||||
key, err := db.Store.GetUserAIProviderKeyByProviderID(ctx, params)
|
||||
if err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
if err := db.decryptUserAIProviderKey(&key); err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) GetUserAIProviderKeysByUserID(ctx context.Context, userID uuid.UUID) ([]database.UserAiProviderKey, error) {
|
||||
keys, err := db.Store.GetUserAIProviderKeysByUserID(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i := range keys {
|
||||
if err := db.decryptUserAIProviderKey(&keys[i]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) GetUserAIProviderKeys(ctx context.Context) ([]database.UserAiProviderKey, error) {
|
||||
keys, err := db.Store.GetUserAIProviderKeys(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i := range keys {
|
||||
if err := db.decryptUserAIProviderKey(&keys[i]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) UpsertUserAIProviderKey(ctx context.Context, params database.UpsertUserAIProviderKeyParams) (database.UserAiProviderKey, error) {
|
||||
if strings.TrimSpace(params.APIKey) == "" {
|
||||
params.ApiKeyKeyID = sql.NullString{}
|
||||
} else if err := db.encryptField(¶ms.APIKey, ¶ms.ApiKeyKeyID); err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
|
||||
key, err := db.Store.UpsertUserAIProviderKey(ctx, params)
|
||||
if err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
if err := db.decryptUserAIProviderKey(&key); err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) UpdateUserAIProviderKey(ctx context.Context, params database.UpdateUserAIProviderKeyParams) (database.UserAiProviderKey, error) {
|
||||
if strings.TrimSpace(params.APIKey) == "" {
|
||||
params.ApiKeyKeyID = sql.NullString{}
|
||||
} else if err := db.encryptField(¶ms.APIKey, ¶ms.ApiKeyKeyID); err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
|
||||
key, err := db.Store.UpdateUserAIProviderKey(ctx, params)
|
||||
if err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
if err := db.decryptUserAIProviderKey(&key); err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) UpdateEncryptedUserAIProviderKey(ctx context.Context, params database.UpdateEncryptedUserAIProviderKeyParams) (database.UserAiProviderKey, error) {
|
||||
if strings.TrimSpace(params.APIKey) == "" {
|
||||
params.ApiKeyKeyID = sql.NullString{}
|
||||
} else if err := db.encryptField(¶ms.APIKey, ¶ms.ApiKeyKeyID); err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
|
||||
key, err := db.Store.UpdateEncryptedUserAIProviderKey(ctx, params)
|
||||
if err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
if err := db.decryptUserAIProviderKey(&key); err != nil {
|
||||
return database.UserAiProviderKey{}, err
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func (db *dbCrypt) decryptUserChatProviderKey(key *database.UserChatProviderKey) error {
|
||||
return db.decryptField(&key.APIKey, key.ApiKeyKeyID)
|
||||
}
|
||||
|
||||
@@ -1292,6 +1292,168 @@ func TestAIProviderKeys(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestUserAIProviderKeys(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
|
||||
const (
|
||||
//nolint:gosec // test credentials
|
||||
initialAPIKey = "sk-initial-ai-provider-key-value"
|
||||
//nolint:gosec // test credentials
|
||||
updatedAPIKey = "sk-updated-ai-provider-key-value"
|
||||
//nolint:gosec // test credentials
|
||||
rotatedAPIKey = "sk-rotated-ai-provider-key-value"
|
||||
)
|
||||
|
||||
insertProviderAndKey := func(
|
||||
t *testing.T,
|
||||
crypt *dbCrypt,
|
||||
ciphers []Cipher,
|
||||
) (database.AIProvider, database.UserAiProviderKey) {
|
||||
t.Helper()
|
||||
user := dbgen.User(t, crypt, database.User{})
|
||||
provider := dbgen.AIProvider(t, crypt, database.AIProvider{})
|
||||
now := dbtime.Now()
|
||||
|
||||
key, err := crypt.UpsertUserAIProviderKey(ctx, database.UpsertUserAIProviderKeyParams{
|
||||
ID: uuid.New(),
|
||||
UserID: user.ID,
|
||||
AIProviderID: provider.ID,
|
||||
APIKey: initialAPIKey,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, initialAPIKey, key.APIKey)
|
||||
require.Equal(t, ciphers[0].HexDigest(), key.ApiKeyKeyID.String)
|
||||
return provider, key
|
||||
}
|
||||
|
||||
getRawUserAIProviderKey := func(t *testing.T, store database.Store, userID uuid.UUID, providerID uuid.UUID) database.UserAiProviderKey {
|
||||
t.Helper()
|
||||
key, err := store.GetUserAIProviderKeyByProviderID(ctx, database.GetUserAIProviderKeyByProviderIDParams{
|
||||
UserID: userID,
|
||||
AIProviderID: providerID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
return key
|
||||
}
|
||||
|
||||
t.Run("UpsertUserAIProviderKeyCreatesValue", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider, key := insertProviderAndKey(t, crypt, ciphers)
|
||||
|
||||
got, err := crypt.GetUserAIProviderKeyByProviderID(ctx, database.GetUserAIProviderKeyByProviderIDParams{
|
||||
UserID: key.UserID,
|
||||
AIProviderID: provider.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, key.ID, got.ID)
|
||||
require.Equal(t, initialAPIKey, got.APIKey)
|
||||
require.Equal(t, ciphers[0].HexDigest(), got.ApiKeyKeyID.String)
|
||||
|
||||
rawKey := getRawUserAIProviderKey(t, db, key.UserID, provider.ID)
|
||||
require.NotEqual(t, initialAPIKey, rawKey.APIKey)
|
||||
requireEncryptedEquals(t, ciphers[0], rawKey.APIKey, initialAPIKey)
|
||||
})
|
||||
|
||||
t.Run("GetUserAIProviderKeysByUserID", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, crypt, ciphers := setup(t)
|
||||
provider, key := insertProviderAndKey(t, crypt, ciphers)
|
||||
|
||||
keys, err := crypt.GetUserAIProviderKeysByUserID(ctx, key.UserID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, keys, 1)
|
||||
require.Equal(t, key.ID, keys[0].ID)
|
||||
require.Equal(t, provider.ID, keys[0].AIProviderID)
|
||||
require.Equal(t, initialAPIKey, keys[0].APIKey)
|
||||
require.Equal(t, ciphers[0].HexDigest(), keys[0].ApiKeyKeyID.String)
|
||||
})
|
||||
|
||||
t.Run("GetUserAIProviderKeys", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, crypt, ciphers := setup(t)
|
||||
provider, key := insertProviderAndKey(t, crypt, ciphers)
|
||||
|
||||
keys, err := crypt.GetUserAIProviderKeys(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, keys, 1)
|
||||
require.Equal(t, key.ID, keys[0].ID)
|
||||
require.Equal(t, key.UserID, keys[0].UserID)
|
||||
require.Equal(t, provider.ID, keys[0].AIProviderID)
|
||||
require.Equal(t, initialAPIKey, keys[0].APIKey)
|
||||
require.Equal(t, ciphers[0].HexDigest(), keys[0].ApiKeyKeyID.String)
|
||||
})
|
||||
|
||||
t.Run("UpsertUserAIProviderKeyUpdatesValue", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider, key := insertProviderAndKey(t, crypt, ciphers)
|
||||
updatedAt := key.UpdatedAt.Add(time.Minute)
|
||||
|
||||
updated, err := crypt.UpsertUserAIProviderKey(ctx, database.UpsertUserAIProviderKeyParams{
|
||||
ID: uuid.New(),
|
||||
UserID: key.UserID,
|
||||
AIProviderID: provider.ID,
|
||||
APIKey: updatedAPIKey,
|
||||
CreatedAt: key.CreatedAt.Add(time.Minute),
|
||||
UpdatedAt: updatedAt,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, key.ID, updated.ID)
|
||||
require.Equal(t, key.CreatedAt, updated.CreatedAt)
|
||||
require.Equal(t, updatedAt, updated.UpdatedAt)
|
||||
require.Equal(t, updatedAPIKey, updated.APIKey)
|
||||
require.Equal(t, ciphers[0].HexDigest(), updated.ApiKeyKeyID.String)
|
||||
|
||||
rawKey := getRawUserAIProviderKey(t, db, key.UserID, provider.ID)
|
||||
require.NotEqual(t, updatedAPIKey, rawKey.APIKey)
|
||||
requireEncryptedEquals(t, ciphers[0], rawKey.APIKey, updatedAPIKey)
|
||||
})
|
||||
|
||||
t.Run("UpdateUserAIProviderKey", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider, key := insertProviderAndKey(t, crypt, ciphers)
|
||||
|
||||
updated, err := crypt.UpdateUserAIProviderKey(ctx, database.UpdateUserAIProviderKeyParams{
|
||||
UserID: key.UserID,
|
||||
AIProviderID: provider.ID,
|
||||
APIKey: updatedAPIKey,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, key.ID, updated.ID)
|
||||
require.WithinDuration(t, dbtime.Now(), updated.UpdatedAt, time.Minute)
|
||||
require.Equal(t, updatedAPIKey, updated.APIKey)
|
||||
require.Equal(t, ciphers[0].HexDigest(), updated.ApiKeyKeyID.String)
|
||||
|
||||
rawKey := getRawUserAIProviderKey(t, db, key.UserID, provider.ID)
|
||||
require.NotEqual(t, updatedAPIKey, rawKey.APIKey)
|
||||
requireEncryptedEquals(t, ciphers[0], rawKey.APIKey, updatedAPIKey)
|
||||
})
|
||||
|
||||
t.Run("UpdateEncryptedUserAIProviderKey", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, crypt, ciphers := setup(t)
|
||||
provider, key := insertProviderAndKey(t, crypt, ciphers)
|
||||
|
||||
updated, err := crypt.UpdateEncryptedUserAIProviderKey(ctx, database.UpdateEncryptedUserAIProviderKeyParams{
|
||||
ID: key.ID,
|
||||
APIKey: rotatedAPIKey,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, key.ID, updated.ID)
|
||||
require.Equal(t, rotatedAPIKey, updated.APIKey)
|
||||
require.Equal(t, ciphers[0].HexDigest(), updated.ApiKeyKeyID.String)
|
||||
|
||||
rawKey := getRawUserAIProviderKey(t, db, key.UserID, provider.ID)
|
||||
require.NotEqual(t, rotatedAPIKey, rawKey.APIKey)
|
||||
requireEncryptedEquals(t, ciphers[0], rawKey.APIKey, rotatedAPIKey)
|
||||
})
|
||||
}
|
||||
|
||||
func TestMCPServerUserTokens(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
|
||||
Reference in New Issue
Block a user