package model import ( "testing" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "gorm.io/gorm" "gorm.io/gorm/logger" ) func setupChannelAssetCredentialDB(t *testing.T) *gorm.DB { t.Helper() db, err := gorm.Open(sqlite.Open("file:channel_asset_credentials?mode=memory&cache=shared"), &gorm.Config{ Logger: logger.Default.LogMode(logger.Silent), }) require.NoError(t, err) sqlDB, err := db.DB() require.NoError(t, err) sqlDB.SetMaxOpenConns(1) originalDB := DB DB = db t.Cleanup(func() { DB = originalDB require.NoError(t, sqlDB.Close()) }) return db } func TestChannelAssetCredentialUpsertKeepsOneCredentialPerChannel(t *testing.T) { db := setupChannelAssetCredentialDB(t) require.NoError(t, db.AutoMigrate(&ChannelAssetCredential{})) require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{ ChannelId: 61, AccessKey: "ak-old", SecretKey: "sk-old", PoolID: "pool-old", })) require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{ ChannelId: 61, AccessKey: "ak-new", SecretKey: "sk-new", PoolID: "pool-new", })) credential, err := GetChannelAssetCredential(61) require.NoError(t, err) require.NotNil(t, credential) assert.Equal(t, "ak-new", credential.AccessKey) assert.Equal(t, "sk-new", credential.SecretKey) assert.Equal(t, "pool-new", credential.PoolID) var count int64 require.NoError(t, db.Model(&ChannelAssetCredential{}).Where("channel_id = ?", 61).Count(&count).Error) assert.Equal(t, int64(1), count) } func TestChannelAssetCredentialSummariesDoNotContainSecrets(t *testing.T) { db := setupChannelAssetCredentialDB(t) require.NoError(t, db.AutoMigrate(&ChannelAssetCredential{})) require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{ ChannelId: 61, AccessKey: "ak-secret", SecretKey: "sk-secret", PoolID: "pool-61", })) summaries, err := GetChannelAssetCredentialSummaries([]int{61, 62}) require.NoError(t, err) require.Contains(t, summaries, 61) assert.Equal(t, "pool-61", summaries[61].PoolID) assert.NotContains(t, summaries, 62) } func TestChannelDeleteRemovesOnlyItsAssetCredential(t *testing.T) { db := setupChannelAssetCredentialDB(t) require.NoError(t, db.AutoMigrate(&Channel{}, &Ability{}, &ChannelAssetCredential{})) for _, id := range []int{61, 62} { require.NoError(t, db.Create(&Channel{Id: id, Key: "key", Name: "channel", Group: "default", Models: "model"}).Error) require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{ChannelId: id, AccessKey: "ak", SecretKey: "sk"})) } require.NoError(t, (&Channel{Id: 61}).Delete()) credential, err := GetChannelAssetCredential(61) require.NoError(t, err) assert.Nil(t, credential) credential, err = GetChannelAssetCredential(62) require.NoError(t, err) assert.NotNil(t, credential) } func TestDeleteDisabledChannelRemovesAssetCredentials(t *testing.T) { db := setupChannelAssetCredentialDB(t) require.NoError(t, db.AutoMigrate(&Channel{}, &Ability{}, &ChannelAssetCredential{})) require.NoError(t, db.Create(&Channel{Id: 61, Key: "key", Name: "disabled", Group: "default", Models: "model", Status: 2}).Error) require.NoError(t, db.Create(&Channel{Id: 62, Key: "key", Name: "enabled", Group: "default", Models: "model", Status: 1}).Error) for _, id := range []int{61, 62} { require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{ChannelId: id, AccessKey: "ak", SecretKey: "sk"})) } _, err := DeleteDisabledChannel() require.NoError(t, err) credential, err := GetChannelAssetCredential(61) require.NoError(t, err) assert.Nil(t, credential) credential, err = GetChannelAssetCredential(62) require.NoError(t, err) assert.NotNil(t, credential) } func TestDeleteChannelByStatusRemovesAssetCredentials(t *testing.T) { db := setupChannelAssetCredentialDB(t) require.NoError(t, db.AutoMigrate(&Channel{}, &Ability{}, &ChannelAssetCredential{})) require.NoError(t, db.Create(&Channel{Id: 61, Key: "key", Name: "disabled", Group: "default", Models: "model", Status: 2}).Error) require.NoError(t, db.Create(&Channel{Id: 62, Key: "key", Name: "enabled", Group: "default", Models: "model", Status: 1}).Error) for _, id := range []int{61, 62} { require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{ChannelId: id, AccessKey: "ak", SecretKey: "sk"})) } _, err := DeleteChannelByStatus(2) require.NoError(t, err) credential, err := GetChannelAssetCredential(61) require.NoError(t, err) assert.Nil(t, credential) credential, err = GetChannelAssetCredential(62) require.NoError(t, err) assert.NotNil(t, credential) }