feat: Implement storage provider management functionality

- Added new routes and handlers for managing storage providers, including creation, editing, and deletion.
- Introduced a new StorageProvider form component for user input.
- Enhanced the database schema to support storage provider references in transfer configurations.
- Implemented encryption for sensitive fields in storage provider data.
- Added tests for storage provider API endpoints and integration with the database.
- Updated frontend components to support storage provider selection and testing.
This commit is contained in:
StarFleetCPTN committed 2025-04-16 17:18:53 -07:00
1 parent 88d0ac815a
commit 31871bd16e
59 files changed
+15842 -594

No files matched your search

+62 -2
View File
@@ -6,12 +6,14 @@ import (
"path/filepath"
"github.com/glebarez/sqlite"
"github.com/starfleetcptn/gomft/internal/db/middleware"
"github.com/starfleetcptn/gomft/internal/db/migrations"
"gorm.io/gorm"
)
type DB struct {
*gorm.DB
encryptionMiddleware *middleware.EncryptionMiddleware
}
func Initialize(dbPath string) (*DB, error) {
@@ -48,7 +50,17 @@ func Initialize(dbPath string) (*DB, error) {
return nil, fmt.Errorf("failed to reconnect to database after migrations: %v", err)
}
return &DB{DB: db}, nil
// Initialize and register the encryption middleware
encryptionMiddleware, err := middleware.NewEncryptionMiddleware()
if err != nil {
return nil, fmt.Errorf("failed to initialize encryption middleware: %v", err)
}
encryptionMiddleware.RegisterHooks(db)
return &DB{
DB: db,
encryptionMiddleware: encryptionMiddleware,
}, nil
}
// ReopenWithoutMigrations reopens the database connection without running migrations
@@ -60,7 +72,17 @@ func ReopenWithoutMigrations(dbPath string) (*DB, error) {
return nil, fmt.Errorf("failed to connect to database: %v", err)
}
return &DB{DB: db}, nil
// Initialize and register the encryption middleware
encryptionMiddleware, err := middleware.NewEncryptionMiddleware()
if err != nil {
return nil, fmt.Errorf("failed to initialize encryption middleware: %v", err)
}
encryptionMiddleware.RegisterHooks(db)
return &DB{
DB: db,
encryptionMiddleware: encryptionMiddleware,
}, nil
}
func (db *DB) Close() error {
@@ -70,3 +92,41 @@ func (db *DB) Close() error {
}
return sqlDB.Close()
}
// EnableEncryption enables the encryption middleware
func (db *DB) EnableEncryption() {
if db.encryptionMiddleware != nil {
db.encryptionMiddleware.Enable()
}
}
// DisableEncryption disables the encryption middleware
func (db *DB) DisableEncryption() {
if db.encryptionMiddleware != nil {
db.encryptionMiddleware.Disable()
}
}
// IsEncryptionEnabled returns whether the encryption middleware is enabled
func (db *DB) IsEncryptionEnabled() bool {
if db.encryptionMiddleware != nil {
return db.encryptionMiddleware.IsEnabled()
}
return false
}
// Connect initializes a database connection using the default path
// This is used by CLI commands to connect to the database
func Connect() (*DB, error) {
// Get data directory from environment or use default
dataDir := os.Getenv("DATA_DIR")
if dataDir == "" {
dataDir = "./data"
}
// Use default database path
dbPath := filepath.Join(dataDir, "gomft.db")
// Initialize the database
return Initialize(dbPath)
}
@@ -0,0 +1,331 @@
package middleware
import (
"errors"
"fmt"
"reflect"
"strings"
"github.com/starfleetcptn/gomft/internal/encryption"
"gorm.io/gorm"
)
// EncryptionMiddleware handles automatic encryption and decryption of model fields
type EncryptionMiddleware struct {
encryptor *encryption.CredentialEncryptor
enabled bool
}
// NewEncryptionMiddleware creates a new middleware instance for encrypting/decrypting fields
func NewEncryptionMiddleware() (*EncryptionMiddleware, error) {
// Get the global credential encryptor
encryptor, err := encryption.GetGlobalCredentialEncryptor()
if err != nil {
return nil, fmt.Errorf("failed to initialize encryption middleware: %w", err)
}
return &EncryptionMiddleware{
encryptor: encryptor,
enabled: true,
}, nil
}
// Enable turns on automatic encryption/decryption
func (m *EncryptionMiddleware) Enable() {
m.enabled = true
}
// Disable turns off automatic encryption/decryption
func (m *EncryptionMiddleware) Disable() {
m.enabled = false
}
// IsEnabled returns whether the middleware is enabled
func (m *EncryptionMiddleware) IsEnabled() bool {
return m.enabled
}
// RegisterHooks registers the encryption/decryption hooks with the GORM instance
func (m *EncryptionMiddleware) RegisterHooks(db *gorm.DB) {
// Register BeforeSave hook to encrypt sensitive fields
db.Callback().Create().Before("gorm:create").Register("encrypt_before_create", m.encryptBeforeSave)
db.Callback().Update().Before("gorm:update").Register("encrypt_before_update", m.encryptBeforeSave)
// Register AfterFind hook to decrypt sensitive fields
db.Callback().Query().After("gorm:after_query").Register("decrypt_after_find", m.decryptAfterFind)
}
// encryptBeforeSave encrypts sensitive fields before saving to the database
func (m *EncryptionMiddleware) encryptBeforeSave(db *gorm.DB) {
if !m.enabled {
return
}
// Get the model value
value := db.Statement.ReflectValue
if value.Kind() == reflect.Ptr {
value = value.Elem()
}
// Skip if the value is not a struct
if value.Kind() != reflect.Struct {
return
}
// Process the model
if err := m.processModelForEncryption(value); err != nil {
db.AddError(fmt.Errorf("encryption middleware error: %w", err))
}
}
// decryptAfterFind decrypts sensitive fields after retrieving from the database
func (m *EncryptionMiddleware) decryptAfterFind(db *gorm.DB) {
if !m.enabled {
return
}
// Get the model value
value := db.Statement.ReflectValue
if value.Kind() == reflect.Ptr {
value = value.Elem()
}
// Handle slice of models
if value.Kind() == reflect.Slice {
for i := 0; i < value.Len(); i++ {
item := value.Index(i)
if item.Kind() == reflect.Ptr {
item = item.Elem()
}
if item.Kind() == reflect.Struct {
if err := m.processModelForDecryption(item); err != nil {
db.AddError(fmt.Errorf("decryption middleware error [index %d]: %w", i, err))
return
}
}
}
return
}
// Skip if the value is not a struct
if value.Kind() != reflect.Struct {
return
}
// Process the model
if err := m.processModelForDecryption(value); err != nil {
db.AddError(fmt.Errorf("decryption middleware error: %w", err))
}
}
// processModelForEncryption encrypts sensitive fields in a model
func (m *EncryptionMiddleware) processModelForEncryption(value reflect.Value) error {
modelType := value.Type()
// Special handling for StorageProvider type
if modelType.Name() == "StorageProvider" {
return m.encryptStorageProvider(value)
}
// Generic handling for models with encryptable fields
for i := 0; i < modelType.NumField(); i++ {
field := modelType.Field(i)
// Check if field requires encryption based on its name
fieldName := field.Name
if requiresEncryption, credType := encryption.RequiresEncryption(fieldName); requiresEncryption {
// Get the field value
fieldValue := value.Field(i)
if !fieldValue.CanInterface() || !fieldValue.CanSet() {
continue
}
// Get string value
strValue, ok := fieldValue.Interface().(string)
if !ok || strValue == "" {
continue
}
// If already encrypted, skip
if m.encryptor.IsEncrypted(strValue) {
continue
}
// Encrypt the field
encryptedValue, err := m.encryptor.Encrypt(strValue, credType)
if err != nil {
return fmt.Errorf("failed to encrypt field %s: %w", fieldName, err)
}
// Find the corresponding encrypted field
encryptedFieldName := "Encrypted" + fieldName
encryptedField := value.FieldByName(encryptedFieldName)
// If encrypted field exists and can be set, set it
if encryptedField.IsValid() && encryptedField.CanSet() {
encryptedField.SetString(encryptedValue)
// If the original field is marked with gorm:"-", we should clear it to prevent leaking it
if field.Tag.Get("gorm") == "-" {
fieldValue.SetString("")
}
}
}
}
return nil
}
// processModelForDecryption decrypts encrypted fields in a model
func (m *EncryptionMiddleware) processModelForDecryption(value reflect.Value) error {
modelType := value.Type()
// Special handling for StorageProvider type
if modelType.Name() == "StorageProvider" {
return m.decryptStorageProvider(value)
}
// Generic handling for models with encrypted fields
for i := 0; i < modelType.NumField(); i++ {
field := modelType.Field(i)
// Look for encrypted fields based on naming pattern
fieldName := field.Name
if strings.HasPrefix(fieldName, "Encrypted") {
originalFieldName := strings.TrimPrefix(fieldName, "Encrypted")
// Get the encrypted field value
encryptedFieldValue := value.Field(i)
if !encryptedFieldValue.CanInterface() {
continue
}
// Get encrypted string value
encryptedValue, ok := encryptedFieldValue.Interface().(string)
if !ok || encryptedValue == "" {
continue
}
// Decrypt the field
decryptedValue, err := m.encryptor.DecryptField(encryptedValue)
if err != nil {
// Log the error but continue
fmt.Printf("Warning: failed to decrypt field %s: %v\n", fieldName, err)
continue
}
// Find the corresponding original field
originalField := value.FieldByName(originalFieldName)
// If original field exists and can be set, set it
if originalField.IsValid() && originalField.CanSet() {
originalField.SetString(decryptedValue)
}
}
}
return nil
}
// encryptStorageProvider handles encryption for StorageProvider model fields
func (m *EncryptionMiddleware) encryptStorageProvider(value reflect.Value) error {
// Check if model implements GetSensitiveFields method
modelInterface := value.Addr().Interface()
// Type assertion to access the GetSensitiveFields method
model, ok := modelInterface.(interface {
GetSensitiveFields() map[string]string
})
if !ok {
return errors.New("StorageProvider model does not implement GetSensitiveFields")
}
// Get sensitive fields that need encryption
sensitiveFields := model.GetSensitiveFields()
// Encrypt each sensitive field
for fieldName, fieldValue := range sensitiveFields {
if fieldValue == "" {
continue
}
// Skip already encrypted values
if m.encryptor.IsEncrypted(fieldValue) {
continue
}
// Determine the credential type based on field name
_, credType := encryption.RequiresEncryption(fieldName)
// Encrypt the value
encryptedValue, err := m.encryptor.Encrypt(fieldValue, credType)
if err != nil {
return fmt.Errorf("failed to encrypt StorageProvider field %s: %w", fieldName, err)
}
// Find the corresponding encrypted field
encryptedFieldName := "Encrypted" + fieldName
encryptedField := value.FieldByName(encryptedFieldName)
// Set the encrypted value
if encryptedField.IsValid() && encryptedField.CanSet() {
encryptedField.SetString(encryptedValue)
// Clear the original field if it shouldn't be stored
originalField := value.FieldByName(fieldName)
if originalField.IsValid() && originalField.CanSet() {
// Find the field in the struct type to check its gorm tag
modelType := reflect.TypeOf(model).Elem()
if field, found := modelType.FieldByName(fieldName); found && field.Tag.Get("gorm") == "-" {
originalField.SetString("")
}
}
}
}
return nil
}
// decryptStorageProvider handles decryption for StorageProvider model fields
func (m *EncryptionMiddleware) decryptStorageProvider(value reflect.Value) error {
// Fields to decrypt
encryptedFields := []string{
"EncryptedPassword",
"EncryptedSecretKey",
"EncryptedClientSecret",
"EncryptedRefreshToken",
}
// Process each encrypted field
for _, fieldName := range encryptedFields {
encryptedField := value.FieldByName(fieldName)
if !encryptedField.IsValid() || !encryptedField.CanInterface() {
continue
}
// Get encrypted value
encryptedValue, ok := encryptedField.Interface().(string)
if !ok || encryptedValue == "" {
continue
}
// Decrypt value
decryptedValue, err := m.encryptor.DecryptField(encryptedValue)
if err != nil {
// Log warning but continue with other fields
fmt.Printf("Warning: failed to decrypt StorageProvider field %s: %v\n", fieldName, err)
continue
}
// Set decrypted value to the original field
originalFieldName := strings.TrimPrefix(fieldName, "Encrypted")
originalField := value.FieldByName(originalFieldName)
if originalField.IsValid() && originalField.CanSet() {
originalField.SetString(decryptedValue)
}
}
return nil
}
@@ -0,0 +1,253 @@
package middleware
import (
"testing"
"time"
"strings"
"github.com/glebarez/sqlite"
"github.com/starfleetcptn/gomft/internal/encryption"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
// TestModel is a simple model for testing encryption middleware
type TestModel struct {
ID uint `gorm:"primarykey"`
Name string `gorm:"not null"`
Password string `gorm:"-"` // Not stored in DB, only for form input
EncryptedPassword string `gorm:"column:encrypted_password"`
APIKey string `gorm:"-"` // Not stored in DB, only for form input
EncryptedAPIKey string `gorm:"column:encrypted_api_key"`
CreatedAt time.Time `gorm:"not null"`
UpdatedAt time.Time `gorm:"not null"`
}
// StorageProvider is a simplified version of the real model for testing
type StorageProvider struct {
ID uint `gorm:"primarykey"`
Name string `gorm:"not null"`
Type string `gorm:"not null"`
Password string `gorm:"-"` // Not stored in DB
EncryptedPassword string `gorm:"column:encrypted_password"`
SecretKey string `gorm:"-"` // Not stored in DB
EncryptedSecretKey string `gorm:"column:encrypted_secret_key"`
ClientSecret string `gorm:"-"` // Not stored in DB
EncryptedClientSecret string `gorm:"column:encrypted_client_secret"`
RefreshToken string `gorm:"-"` // Not stored in DB
EncryptedRefreshToken string `gorm:"column:encrypted_refresh_token"`
CreatedAt time.Time `gorm:"not null"`
UpdatedAt time.Time `gorm:"not null"`
}
// GetSensitiveFields returns a map of field names to values that need encryption
func (sp *StorageProvider) GetSensitiveFields() map[string]string {
sensitiveFields := make(map[string]string)
if sp.Password != "" {
sensitiveFields["Password"] = sp.Password
}
if sp.SecretKey != "" {
sensitiveFields["SecretKey"] = sp.SecretKey
}
if sp.ClientSecret != "" {
sensitiveFields["ClientSecret"] = sp.ClientSecret
}
if sp.RefreshToken != "" {
sensitiveFields["RefreshToken"] = sp.RefreshToken
}
return sensitiveFields
}
func setupTestDB(t *testing.T) *gorm.DB {
// Initialize in-memory SQLite database
db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
require.NoError(t, err, "Failed to connect to in-memory database")
// Migrate the test models
err = db.AutoMigrate(&TestModel{}, &StorageProvider{})
require.NoError(t, err, "Failed to migrate test models")
return db
}
func setupEncryptionMiddleware(t *testing.T) (*EncryptionMiddleware, error) {
// Initialize the encryption key manager for testing
err := encryption.InitializeKeyManager("test-key")
require.NoError(t, err, "Failed to initialize key manager")
// Create the encryption middleware
return NewEncryptionMiddleware()
}
func TestEncryptionMiddlewareWithGenericModel(t *testing.T) {
// Setup
db := setupTestDB(t)
middleware, err := setupEncryptionMiddleware(t)
require.NoError(t, err, "Failed to setup encryption middleware")
// Register hooks with GORM
middleware.RegisterHooks(db)
// Create a test model
testModel := &TestModel{
Name: "Test User",
Password: "securePassword123",
APIKey: "api-key-12345",
}
// Save the model - should trigger encryption
err = db.Create(testModel).Error
require.NoError(t, err, "Failed to save test model")
// Verify encrypted fields are set and original fields are cleared
assert.Empty(t, testModel.Password, "Password should be cleared after save")
assert.Empty(t, testModel.APIKey, "APIKey should be cleared after save")
assert.NotEmpty(t, testModel.EncryptedPassword, "EncryptedPassword should be set")
assert.NotEmpty(t, testModel.EncryptedAPIKey, "EncryptedAPIKey should be set")
assert.True(t, strings.HasPrefix(testModel.EncryptedPassword, encryption.EncryptedPrefix), "EncryptedPassword should have encryption prefix")
assert.True(t, strings.HasPrefix(testModel.EncryptedAPIKey, encryption.EncryptedPrefix), "EncryptedAPIKey should have encryption prefix")
// Test retrieval and automatic decryption
retrievedModel := new(TestModel)
err = db.First(retrievedModel, testModel.ID).Error
require.NoError(t, err, "Failed to retrieve test model")
// Verify decryption
assert.Equal(t, "securePassword123", retrievedModel.Password, "Password should be automatically decrypted")
assert.Equal(t, "api-key-12345", retrievedModel.APIKey, "APIKey should be automatically decrypted")
assert.NotEmpty(t, retrievedModel.EncryptedPassword, "EncryptedPassword should remain set")
assert.NotEmpty(t, retrievedModel.EncryptedAPIKey, "EncryptedAPIKey should remain set")
}
func TestEncryptionMiddlewareWithStorageProvider(t *testing.T) {
// Setup
db := setupTestDB(t)
middleware, err := setupEncryptionMiddleware(t)
require.NoError(t, err, "Failed to setup encryption middleware")
// Register hooks with GORM
middleware.RegisterHooks(db)
// Create a storage provider
provider := &StorageProvider{
Name: "Test S3",
Type: "s3",
Password: "testPassword",
SecretKey: "testSecretKey",
ClientSecret: "testClientSecret",
RefreshToken: "testRefreshToken",
}
// Save the provider - should trigger encryption
err = db.Create(provider).Error
require.NoError(t, err, "Failed to save storage provider")
// Verify encrypted fields are set and original fields are cleared
assert.Empty(t, provider.Password, "Password should be cleared after save")
assert.Empty(t, provider.SecretKey, "SecretKey should be cleared after save")
assert.Empty(t, provider.ClientSecret, "ClientSecret should be cleared after save")
assert.Empty(t, provider.RefreshToken, "RefreshToken should be cleared after save")
assert.NotEmpty(t, provider.EncryptedPassword, "EncryptedPassword should be set")
assert.NotEmpty(t, provider.EncryptedSecretKey, "EncryptedSecretKey should be set")
assert.NotEmpty(t, provider.EncryptedClientSecret, "EncryptedClientSecret should be set")
assert.NotEmpty(t, provider.EncryptedRefreshToken, "EncryptedRefreshToken should be set")
// Test retrieval and automatic decryption
retrievedProvider := new(StorageProvider)
err = db.First(retrievedProvider, provider.ID).Error
require.NoError(t, err, "Failed to retrieve storage provider")
// Verify decryption
assert.Equal(t, "testPassword", retrievedProvider.Password, "Password should be automatically decrypted")
assert.Equal(t, "testSecretKey", retrievedProvider.SecretKey, "SecretKey should be automatically decrypted")
assert.Equal(t, "testClientSecret", retrievedProvider.ClientSecret, "ClientSecret should be automatically decrypted")
assert.Equal(t, "testRefreshToken", retrievedProvider.RefreshToken, "RefreshToken should be automatically decrypted")
}
func TestEncryptionMiddlewareWithMultipleRecords(t *testing.T) {
// Setup
db := setupTestDB(t)
middleware, err := setupEncryptionMiddleware(t)
require.NoError(t, err, "Failed to setup encryption middleware")
// Register hooks with GORM
middleware.RegisterHooks(db)
// Create multiple test models
models := []TestModel{
{Name: "User 1", Password: "password1", APIKey: "apikey1"},
{Name: "User 2", Password: "password2", APIKey: "apikey2"},
{Name: "User 3", Password: "password3", APIKey: "apikey3"},
}
// Save all models
err = db.Create(&models).Error
require.NoError(t, err, "Failed to save multiple test models")
// Retrieve all models
var retrievedModels []TestModel
err = db.Find(&retrievedModels).Error
require.NoError(t, err, "Failed to retrieve all test models")
// Verify count
assert.Equal(t, 3, len(retrievedModels), "Should retrieve 3 models")
// Verify each model was properly decrypted
expectedPasswords := []string{"password1", "password2", "password3"}
expectedAPIKeys := []string{"apikey1", "apikey2", "apikey3"}
for i, model := range retrievedModels {
assert.Equal(t, expectedPasswords[i], model.Password, "Password should be automatically decrypted")
assert.Equal(t, expectedAPIKeys[i], model.APIKey, "APIKey should be automatically decrypted")
assert.NotEmpty(t, model.EncryptedPassword, "EncryptedPassword should remain set")
assert.NotEmpty(t, model.EncryptedAPIKey, "EncryptedAPIKey should remain set")
}
}
func TestEncryptionMiddlewareDisabled(t *testing.T) {
// Setup
db := setupTestDB(t)
middleware, err := setupEncryptionMiddleware(t)
require.NoError(t, err, "Failed to setup encryption middleware")
// Register hooks with GORM
middleware.RegisterHooks(db)
// Disable the middleware
middleware.Disable()
assert.False(t, middleware.IsEnabled(), "Middleware should be disabled")
// Create a test model
testModel := &TestModel{
Name: "Test User",
Password: "securePassword123",
APIKey: "api-key-12345",
}
// Save the model - should NOT trigger encryption since middleware is disabled
err = db.Create(testModel).Error
require.NoError(t, err, "Failed to save test model")
// Verify sensitive fields are NOT encrypted
assert.Equal(t, "securePassword123", testModel.Password, "Password should not be cleared when middleware is disabled")
assert.Equal(t, "api-key-12345", testModel.APIKey, "APIKey should not be cleared when middleware is disabled")
assert.Empty(t, testModel.EncryptedPassword, "EncryptedPassword should not be set when middleware is disabled")
assert.Empty(t, testModel.EncryptedAPIKey, "EncryptedAPIKey should not be set when middleware is disabled")
// Re-enable the middleware for subsequent operations
middleware.Enable()
assert.True(t, middleware.IsEnabled(), "Middleware should be enabled")
// Update the model - should now trigger encryption
testModel.Password = "newPassword456"
err = db.Save(testModel).Error
require.NoError(t, err, "Failed to update test model")
// Verify encryption now happened
assert.Empty(t, testModel.Password, "Password should be cleared after update with middleware enabled")
assert.NotEmpty(t, testModel.EncryptedPassword, "EncryptedPassword should be set after update with middleware enabled")
}
+922
View File
@@ -0,0 +1,922 @@
package db
import (
"fmt"
"log"
"strings"
"time"
"github.com/starfleetcptn/gomft/internal/encryption"
)
// ProviderConfig represents a unique provider configuration extracted from TransferConfigs
type ProviderConfig struct {
// Common identification fields
SourceOrDest string // "source" or "destination"
Type StorageProviderType // The provider type
// All possible provider fields
Host string
Port int
Username string
Password string
KeyFile string
Bucket string
Region string
AccessKey string
SecretKey string
Endpoint string
Share string
Domain string
PassiveMode *bool
ClientID string
ClientSecret string
DriveID string
TeamDrive string
ReadOnly *bool
StartYear int
IncludeArchived *bool
UseBuiltinAuth *bool
// Status
Authenticated *bool
// References
ConfigIDs []uint // IDs of TransferConfigs using this provider config
CreatedBy uint // User ID who created the config
// For mapping to created provider
NewProviderID uint // ID of the created StorageProvider (used during migration)
}
// GetUniqueKey returns a string key that uniquely identifies this provider configuration
// This is used for deduplication
func (pc *ProviderConfig) GetUniqueKey() string {
// Create a composite key based on the most important identifying fields
// The combination of fields depends on the provider type
switch pc.Type {
case ProviderTypeSFTP, ProviderTypeHetzner, ProviderTypeFTP:
return fmt.Sprintf("%s:%s:%d:%s:%s",
pc.Type, pc.Host, pc.Port, pc.Username, pc.KeyFile)
case ProviderTypeS3:
return fmt.Sprintf("%s:%s:%s:%s",
pc.Type, pc.Endpoint, pc.Region, pc.AccessKey)
case ProviderTypeSMB:
return fmt.Sprintf("%s:%s:%s:%s",
pc.Type, pc.Host, pc.Share, pc.Username)
case ProviderTypeOneDrive, ProviderTypeGoogleDrive, ProviderTypeGooglePhoto:
return fmt.Sprintf("%s:%s:%s",
pc.Type, pc.ClientID, pc.DriveID)
case ProviderTypeLocal:
return fmt.Sprintf("%s:%d", pc.Type, pc.CreatedBy)
default:
// Fallback for unknown types
return fmt.Sprintf("%s:%s:%d:%s",
pc.Type, pc.Host, pc.Port, pc.Username)
}
}
// GenerateName generates a meaningful name for the provider
func (pc *ProviderConfig) GenerateName(configName string) string {
if configName == "" {
configName = "Unnamed Config"
}
basePrefix := ""
if pc.SourceOrDest == "source" {
basePrefix = "Source -"
} else {
basePrefix = "Destination -"
}
// Include identifiable information based on provider type
switch pc.Type {
case ProviderTypeSFTP, ProviderTypeHetzner, ProviderTypeFTP:
return fmt.Sprintf("%s %s %s (%s@%s)", configName, basePrefix, pc.Type, pc.Username, pc.Host)
case ProviderTypeS3:
return fmt.Sprintf("%s %s %s (%s - %s)", configName, basePrefix, pc.Type, pc.Region, pc.Bucket)
case ProviderTypeSMB:
return fmt.Sprintf("%s %s %s (%s on %s)", configName, basePrefix, pc.Type, pc.Share, pc.Host)
case ProviderTypeOneDrive:
return fmt.Sprintf("%s %s OneDrive", configName, basePrefix)
case ProviderTypeGoogleDrive:
return fmt.Sprintf("%s %s Google Drive", configName, basePrefix)
case ProviderTypeGooglePhoto:
return fmt.Sprintf("%s %s Google Photos", configName, basePrefix)
case ProviderTypeLocal:
return fmt.Sprintf("%s %s Local", configName, basePrefix)
default:
return fmt.Sprintf("%s %s %s", configName, basePrefix, pc.Type)
}
}
// MigrationStats holds statistics about the migration process
type MigrationStats struct {
TotalConfigs int
UniqueSourceProviders int
UniqueDestinationProviders int
NewProvidersCreated int
ConfigsUpdated int
Errors []string
StartTime time.Time
EndTime time.Time
}
// MigrationBackup holds backup data for rollback in case of migration failure
type MigrationBackup struct {
Configs []TransferConfig
ProvidersCreated []uint
}
// MigrateProviderDataOptions contains options for the migration process
type MigrateProviderDataOptions struct {
DryRun bool // If true, perform a simulation without actually modifying data
ValidationOnly bool // If true, only perform validation without migration
Force bool // If true, ignore validation errors and proceed with migration
BackupDir string // Directory to store backups in
}
// ExtractUniqueProviderConfigs extracts all unique provider configurations from existing TransferConfig records
// It returns a map of provider keys to ProviderConfig objects and any error encountered
func (db *DB) ExtractUniqueProviderConfigs() (map[string]*ProviderConfig, error) {
log.Println("Starting extraction of unique provider configurations...")
// Get all transfer configs
var configs []TransferConfig
if err := db.Find(&configs).Error; err != nil {
return nil, fmt.Errorf("failed to retrieve transfer configs: %v", err)
}
log.Printf("Found %d transfer configs", len(configs))
// Map to store unique provider configurations
uniqueProviders := make(map[string]*ProviderConfig)
// Process each transfer config
for _, config := range configs {
// Skip if already using provider references
if config.IsUsingProviderReferences() {
log.Printf("Config ID %d already using provider references, skipping", config.ID)
continue
}
// Process source provider if not already using a reference
if !config.IsUsingSourceProviderReference() && config.SourceType != "" {
sourceConfig := extractSourceProviderConfig(&config)
key := sourceConfig.GetUniqueKey()
if existing, exists := uniqueProviders[key]; exists {
// Add this config ID to the existing provider's references
existing.ConfigIDs = append(existing.ConfigIDs, config.ID)
log.Printf("Added Config ID %d to existing source provider key %s", config.ID, key)
} else {
// Add this as a new unique provider
uniqueProviders[key] = sourceConfig
log.Printf("Added new unique source provider with key %s", key)
}
}
// Process destination provider if not already using a reference
if !config.IsUsingDestinationProviderReference() && config.DestinationType != "" {
destConfig := extractDestinationProviderConfig(&config)
key := destConfig.GetUniqueKey()
if existing, exists := uniqueProviders[key]; exists {
// Add this config ID to the existing provider's references
existing.ConfigIDs = append(existing.ConfigIDs, config.ID)
log.Printf("Added Config ID %d to existing destination provider key %s", config.ID, key)
} else {
// Add this as a new unique provider
uniqueProviders[key] = destConfig
log.Printf("Added new unique destination provider with key %s", key)
}
}
}
// Count the number of source and destination providers
sourceCount := 0
destCount := 0
for _, provider := range uniqueProviders {
if provider.SourceOrDest == "source" {
sourceCount++
} else {
destCount++
}
}
log.Printf("Extraction complete. Found %d unique provider configurations (%d source, %d destination)",
len(uniqueProviders), sourceCount, destCount)
return uniqueProviders, nil
}
// extractSourceProviderConfig extracts source provider details from a TransferConfig
func extractSourceProviderConfig(config *TransferConfig) *ProviderConfig {
sourceConfig := &ProviderConfig{
SourceOrDest: "source",
Type: StorageProviderType(config.SourceType),
CreatedBy: config.CreatedBy,
ConfigIDs: []uint{config.ID},
// Copy all relevant source fields
Host: config.SourceHost,
Port: config.SourcePort,
Username: config.SourceUser,
Password: config.SourcePassword,
KeyFile: config.SourceKeyFile,
Bucket: config.SourceBucket,
Region: config.SourceRegion,
AccessKey: config.SourceAccessKey,
SecretKey: config.SourceSecretKey,
Endpoint: config.SourceEndpoint,
Share: config.SourceShare,
Domain: config.SourceDomain,
PassiveMode: config.SourcePassiveMode,
ClientID: config.SourceClientID,
ClientSecret: config.SourceClientSecret,
DriveID: config.SourceDriveID,
TeamDrive: config.SourceTeamDrive,
UseBuiltinAuth: config.UseBuiltinAuthSource,
}
// Handle boolean pointers
if config.SourceReadOnly != nil {
sourceConfig.ReadOnly = config.SourceReadOnly
}
if config.SourceIncludeArchived != nil {
sourceConfig.IncludeArchived = config.SourceIncludeArchived
}
// Special handling for OAuth authentication status
if config.SourceType == "gdrive" || config.SourceType == "gphotos" {
authenticated := config.GetGoogleAuthenticated()
sourceConfig.Authenticated = &authenticated
}
sourceConfig.StartYear = config.SourceStartYear
return sourceConfig
}
// extractDestinationProviderConfig extracts destination provider details from a TransferConfig
func extractDestinationProviderConfig(config *TransferConfig) *ProviderConfig {
destConfig := &ProviderConfig{
SourceOrDest: "destination",
Type: StorageProviderType(config.DestinationType),
CreatedBy: config.CreatedBy,
ConfigIDs: []uint{config.ID},
// Copy all relevant destination fields
Host: config.DestHost,
Port: config.DestPort,
Username: config.DestUser,
Password: config.DestPassword,
KeyFile: config.DestKeyFile,
Bucket: config.DestBucket,
Region: config.DestRegion,
AccessKey: config.DestAccessKey,
SecretKey: config.DestSecretKey,
Endpoint: config.DestEndpoint,
Share: config.DestShare,
Domain: config.DestDomain,
PassiveMode: config.DestPassiveMode,
ClientID: config.DestClientID,
ClientSecret: config.DestClientSecret,
DriveID: config.DestDriveID,
TeamDrive: config.DestTeamDrive,
UseBuiltinAuth: config.UseBuiltinAuthDest,
}
// Handle boolean pointers
if config.DestReadOnly != nil {
destConfig.ReadOnly = config.DestReadOnly
}
if config.DestIncludeArchived != nil {
destConfig.IncludeArchived = config.DestIncludeArchived
}
// Special handling for OAuth authentication status
if config.DestinationType == "gdrive" || config.DestinationType == "gphotos" {
authenticated := config.GetGoogleAuthenticated()
destConfig.Authenticated = &authenticated
}
destConfig.StartYear = config.DestStartYear
return destConfig
}
// CreateStorageProviderRecords creates new StorageProvider records from unique provider configurations
// It returns a map of provider keys to new StorageProvider IDs and any error encountered
func (db *DB) CreateStorageProviderRecords(uniqueConfigs map[string]*ProviderConfig) (map[string]uint, error) {
log.Println("Starting creation of StorageProvider records...")
// Get the credential encryptor
credentialEncryptor, err := encryption.GetGlobalCredentialEncryptor()
if err != nil {
return nil, fmt.Errorf("failed to get credential encryptor: %v", err)
}
// Map to store provider keys to their IDs
providerIDMap := make(map[string]uint)
// Start transaction
tx := db.Begin()
if tx.Error != nil {
return nil, fmt.Errorf("failed to start transaction: %v", tx.Error)
}
// Create a function to handle rollback in case of error
rollback := func(err error) (map[string]uint, error) {
tx.Rollback()
return nil, err
}
// Process each unique provider config
for key, providerConfig := range uniqueConfigs {
log.Printf("Creating StorageProvider for key %s...", key)
// Get the first config ID for naming
var configName string
if len(providerConfig.ConfigIDs) > 0 {
firstConfigID := providerConfig.ConfigIDs[0]
var config TransferConfig
if err := tx.First(&config, firstConfigID).Error; err == nil {
configName = config.Name
}
}
// Create new StorageProvider record
provider := &StorageProvider{
Name: providerConfig.GenerateName(configName),
Type: providerConfig.Type,
Host: providerConfig.Host,
Port: providerConfig.Port,
Username: providerConfig.Username,
KeyFile: providerConfig.KeyFile,
Bucket: providerConfig.Bucket,
Region: providerConfig.Region,
AccessKey: providerConfig.AccessKey,
Endpoint: providerConfig.Endpoint,
Share: providerConfig.Share,
Domain: providerConfig.Domain,
PassiveMode: providerConfig.PassiveMode,
ClientID: providerConfig.ClientID,
DriveID: providerConfig.DriveID,
TeamDrive: providerConfig.TeamDrive,
ReadOnly: providerConfig.ReadOnly,
StartYear: providerConfig.StartYear,
IncludeArchived: providerConfig.IncludeArchived,
UseBuiltinAuth: providerConfig.UseBuiltinAuth,
Authenticated: providerConfig.Authenticated,
CreatedBy: providerConfig.CreatedBy,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Encrypt sensitive fields
if providerConfig.Password != "" {
encryptedPwd, err := credentialEncryptor.EncryptPassword(providerConfig.Password)
if err != nil {
return rollback(fmt.Errorf("failed to encrypt password: %v", err))
}
provider.EncryptedPassword = encryptedPwd
}
if providerConfig.SecretKey != "" {
encryptedSecret, err := credentialEncryptor.EncryptSecretKey(providerConfig.SecretKey)
if err != nil {
return rollback(fmt.Errorf("failed to encrypt secret key: %v", err))
}
provider.EncryptedSecretKey = encryptedSecret
}
if providerConfig.ClientSecret != "" {
encryptedClientSecret, err := credentialEncryptor.EncryptField(providerConfig.ClientSecret, encryption.TypeGeneric)
if err != nil {
return rollback(fmt.Errorf("failed to encrypt client secret: %v", err))
}
provider.EncryptedClientSecret = encryptedClientSecret
}
// Create the provider record
if err := tx.Create(provider).Error; err != nil {
sanitizedErrMsg := sanitizeErrorMessage(err.Error())
return rollback(fmt.Errorf("failed to create provider record: %s", sanitizedErrMsg))
}
// Store the provider ID in the map
providerIDMap[key] = provider.ID
// Update the provider config with the new ID
providerConfig.NewProviderID = provider.ID
log.Printf("Created StorageProvider ID %d for key %s", provider.ID, key)
}
// Commit the transaction
if err := tx.Commit().Error; err != nil {
return nil, fmt.Errorf("failed to commit transaction: %v", err)
}
log.Printf("Successfully created %d StorageProvider records", len(providerIDMap))
return providerIDMap, nil
}
// sanitizeErrorMessage removes any potential sensitive information from error messages
func sanitizeErrorMessage(errMsg string) string {
// List of sensitive keywords to check for
sensitiveKeywords := []string{
"password", "secret", "token", "key", "credential", "auth",
}
// Check if the error message contains sensitive information
lowercaseMsg := strings.ToLower(errMsg)
for _, keyword := range sensitiveKeywords {
if strings.Contains(lowercaseMsg, keyword) {
// If it contains sensitive info, return a generic message
return "database error (details omitted for security)"
}
}
return errMsg
}
// UpdateTransferConfigReferences updates TransferConfig records to reference the newly created StorageProvider entities
func (db *DB) UpdateTransferConfigReferences(uniqueConfigs map[string]*ProviderConfig) error {
log.Println("Starting update of TransferConfig references...")
// Start transaction
tx := db.Begin()
if tx.Error != nil {
return fmt.Errorf("failed to start transaction: %v", tx.Error)
}
// Create a function to handle rollback in case of error
rollback := func(err error) error {
tx.Rollback()
return err
}
// Create a map of config IDs to their updates
// This helps batch config updates by config ID
configUpdates := make(map[uint]struct {
SourceProviderID *uint
DestinationProviderID *uint
})
// Process each unique provider config
for _, providerConfig := range uniqueConfigs {
// Skip if no provider ID was assigned (shouldn't happen)
if providerConfig.NewProviderID == 0 {
log.Printf("Warning: Provider config %s has no ID assigned, skipping", providerConfig.GetUniqueKey())
continue
}
// For each config ID that uses this provider
for _, configID := range providerConfig.ConfigIDs {
// Get or initialize the update record
update, exists := configUpdates[configID]
if !exists {
update = struct {
SourceProviderID *uint
DestinationProviderID *uint
}{nil, nil}
}
// Update the appropriate provider ID
if providerConfig.SourceOrDest == "source" {
newID := providerConfig.NewProviderID
update.SourceProviderID = &newID
} else {
newID := providerConfig.NewProviderID
update.DestinationProviderID = &newID
}
// Store the update
configUpdates[configID] = update
}
}
// Apply the updates
totalUpdated := 0
for configID, update := range configUpdates {
// Retrieve the config
var config TransferConfig
if err := tx.First(&config, configID).Error; err != nil {
return rollback(fmt.Errorf("failed to retrieve config ID %d: %v", configID, err))
}
// Update source provider reference if needed
if update.SourceProviderID != nil {
config.SourceProviderID = update.SourceProviderID
}
// Update destination provider reference if needed
if update.DestinationProviderID != nil {
config.DestinationProviderID = update.DestinationProviderID
}
// Save the updated config
if err := tx.Save(&config).Error; err != nil {
return rollback(fmt.Errorf("failed to update config ID %d: %v", configID, err))
}
log.Printf("Updated TransferConfig ID %d with provider references", configID)
totalUpdated++
}
// Commit the transaction
if err := tx.Commit().Error; err != nil {
return fmt.Errorf("failed to commit transaction: %v", err)
}
log.Printf("Successfully updated %d TransferConfig records with provider references", totalUpdated)
return nil
}
// ValidationResult represents the result of a migration validation
type ValidationResult struct {
Success bool
TotalConfigs int
ValidConfigs int
InvalidConfigs int
MissingProviders int
ValidationErrors []string
ConfigsWithErrors []uint
}
// ValidateMigrationIntegrity validates the integrity of the migration
func (db *DB) ValidateMigrationIntegrity() (*ValidationResult, error) {
log.Println("Starting validation of migration integrity...")
result := &ValidationResult{
Success: true,
ValidationErrors: []string{},
ConfigsWithErrors: []uint{},
}
// Get all transfer configs
var configs []TransferConfig
if err := db.Preload("SourceProvider").Preload("DestinationProvider").Find(&configs).Error; err != nil {
return nil, fmt.Errorf("failed to retrieve transfer configs: %v", err)
}
result.TotalConfigs = len(configs)
log.Printf("Found %d transfer configs for validation", result.TotalConfigs)
// Get the credential encryptor for testing decryption
credentialEncryptor, err := encryption.GetGlobalCredentialEncryptor()
if err != nil {
return nil, fmt.Errorf("failed to get credential encryptor: %v", err)
}
// Validate each config
for _, config := range configs {
configValid := true
// Check if this config should be using provider references
shouldUseProviders := !strings.HasPrefix(config.SourceType, "local") || !strings.HasPrefix(config.DestinationType, "local")
// If it should be using provider references but isn't, mark as invalid
if shouldUseProviders && !config.IsUsingProviderReferences() {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d is not using provider references", config.ID))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
}
// Check source provider reference if needed
if !strings.HasPrefix(config.SourceType, "local") && !config.IsUsingSourceProviderReference() {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d is missing source provider reference", config.ID))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
}
// Check destination provider reference if needed
if !strings.HasPrefix(config.DestinationType, "local") && !config.IsUsingDestinationProviderReference() {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d is missing destination provider reference", config.ID))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
}
// If using source provider reference, validate the provider
if config.IsUsingSourceProviderReference() {
if config.SourceProvider == nil {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d has source provider reference but provider is nil", config.ID))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
result.MissingProviders++
} else if config.SourceProvider.Type != StorageProviderType(config.SourceType) {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d source provider type mismatch: config=%s, provider=%s",
config.ID, config.SourceType, config.SourceProvider.Type))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
}
}
// If using destination provider reference, validate the provider
if config.IsUsingDestinationProviderReference() {
if config.DestinationProvider == nil {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d has destination provider reference but provider is nil", config.ID))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
result.MissingProviders++
} else if config.DestinationProvider.Type != StorageProviderType(config.DestinationType) {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d destination provider type mismatch: config=%s, provider=%s",
config.ID, config.DestinationType, config.DestinationProvider.Type))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
}
}
// Verify that source credentials can be retrieved
if !strings.HasPrefix(config.SourceType, "local") {
sourceCreds, err := config.GetSourceCredentials(db)
if err != nil {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d failed to get source credentials: %v", config.ID, err))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
} else {
// Check if credentials are properly encrypted
if encPwd, ok := sourceCreds["encrypted_password"].(string); ok && encPwd != "" {
if !credentialEncryptor.IsEncrypted(encPwd) {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d source password is not properly encrypted", config.ID))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
}
}
if encSecret, ok := sourceCreds["encrypted_secret_key"].(string); ok && encSecret != "" {
if !credentialEncryptor.IsEncrypted(encSecret) {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d source secret key is not properly encrypted", config.ID))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
}
}
}
}
// Verify that destination credentials can be retrieved
if !strings.HasPrefix(config.DestinationType, "local") {
destCreds, err := config.GetDestinationCredentials(db)
if err != nil {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d failed to get destination credentials: %v", config.ID, err))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
} else {
// Check if credentials are properly encrypted
if encPwd, ok := destCreds["encrypted_password"].(string); ok && encPwd != "" {
if !credentialEncryptor.IsEncrypted(encPwd) {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d destination password is not properly encrypted", config.ID))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
}
}
if encSecret, ok := destCreds["encrypted_secret_key"].(string); ok && encSecret != "" {
if !credentialEncryptor.IsEncrypted(encSecret) {
result.ValidationErrors = append(result.ValidationErrors,
fmt.Sprintf("Config ID %d destination secret key is not properly encrypted", config.ID))
result.ConfigsWithErrors = append(result.ConfigsWithErrors, config.ID)
configValid = false
}
}
}
}
if configValid {
result.ValidConfigs++
} else {
result.InvalidConfigs++
result.Success = false
}
}
// Log validation summary
if result.Success {
log.Printf("Validation successful. All %d configs are valid.", result.ValidConfigs)
} else {
log.Printf("Validation failed. %d valid configs, %d invalid configs, %d missing providers.",
result.ValidConfigs, result.InvalidConfigs, result.MissingProviders)
}
return result, nil
}
// MigrateProviderData is the main function that performs the complete migration process
func (db *DB) MigrateProviderData(options MigrateProviderDataOptions) (*MigrationStats, error) {
// Initialize migration stats
stats := &MigrationStats{
StartTime: time.Now(),
Errors: []string{},
}
log.Println("Starting provider data migration...")
// Create backup if not in dry run mode
var backup *MigrationBackup
var err error
if !options.DryRun {
backup, err = db.createMigrationBackup(options.BackupDir)
if err != nil {
return stats, fmt.Errorf("failed to create backup: %v", err)
}
log.Println("Created migration backup")
}
// Extract unique provider configs
uniqueConfigs, err := db.ExtractUniqueProviderConfigs()
if err != nil {
return stats, fmt.Errorf("failed to extract unique provider configurations: %v", err)
}
// Count the number of source and destination providers
sourceCount := 0
destCount := 0
for _, provider := range uniqueConfigs {
if provider.SourceOrDest == "source" {
sourceCount++
} else {
destCount++
}
}
stats.TotalConfigs = len(uniqueConfigs)
stats.UniqueSourceProviders = sourceCount
stats.UniqueDestinationProviders = destCount
// Return if validation only mode
if options.ValidationOnly {
log.Println("Validation-only mode: Migration stopped after extraction")
stats.EndTime = time.Now()
return stats, nil
}
// Return if dry run mode
if options.DryRun {
log.Println("Dry run mode: Migration stopped after extraction")
stats.EndTime = time.Now()
return stats, nil
}
// Create provider records
providerIDMap, err := db.CreateStorageProviderRecords(uniqueConfigs)
if err != nil {
// Attempt rollback
if rollbackErr := db.rollbackMigration(backup); rollbackErr != nil {
stats.Errors = append(stats.Errors, fmt.Sprintf("failed to rollback after provider creation error: %v", rollbackErr))
}
return stats, fmt.Errorf("failed to create provider records: %v", err)
}
stats.NewProvidersCreated = len(providerIDMap)
// Update config references
if err := db.UpdateTransferConfigReferences(uniqueConfigs); err != nil {
// Attempt rollback
if rollbackErr := db.rollbackMigration(backup); rollbackErr != nil {
stats.Errors = append(stats.Errors, fmt.Sprintf("failed to rollback after reference update error: %v", rollbackErr))
}
return stats, fmt.Errorf("failed to update config references: %v", err)
}
// Validate the migration
validationResult, err := db.ValidateMigrationIntegrity()
if err != nil {
stats.Errors = append(stats.Errors, fmt.Sprintf("validation error: %v", err))
// Don't rollback here since the migration might be fine even if validation had errors
}
if validationResult != nil {
stats.ConfigsUpdated = validationResult.ValidConfigs
// If validation failed and not in force mode, rollback
if !validationResult.Success && !options.Force {
log.Println("Validation failed and not in force mode, rolling back...")
if rollbackErr := db.rollbackMigration(backup); rollbackErr != nil {
stats.Errors = append(stats.Errors, fmt.Sprintf("failed to rollback after validation failure: %v", rollbackErr))
}
stats.Errors = append(stats.Errors, validationResult.ValidationErrors...)
return stats, fmt.Errorf("migration validation failed")
}
// If validation failed but in force mode, log warnings
if !validationResult.Success && options.Force {
log.Println("Validation failed but running in force mode, proceeding anyway...")
stats.Errors = append(stats.Errors, "Migration had validation errors but continued due to force mode")
stats.Errors = append(stats.Errors, validationResult.ValidationErrors...)
}
}
stats.EndTime = time.Now()
log.Printf("Migration completed in %v", stats.EndTime.Sub(stats.StartTime))
return stats, nil
}
// createMigrationBackup creates a backup of the current state for rollback
func (db *DB) createMigrationBackup(backupDir string) (*MigrationBackup, error) {
backup := &MigrationBackup{
Configs: []TransferConfig{},
ProvidersCreated: []uint{},
}
// Get all transfer configs
if err := db.Find(&backup.Configs).Error; err != nil {
return nil, fmt.Errorf("failed to backup transfer configs: %v", err)
}
log.Printf("Backed up %d transfer config records", len(backup.Configs))
return backup, nil
}
// rollbackMigration restores the system to its pre-migration state
func (db *DB) rollbackMigration(backup *MigrationBackup) error {
if backup == nil {
return fmt.Errorf("cannot rollback: no backup provided")
}
log.Println("Starting migration rollback...")
// Start transaction
tx := db.Begin()
if tx.Error != nil {
return fmt.Errorf("failed to start rollback transaction: %v", tx.Error)
}
// First, delete any provider records created during the migration
if len(backup.ProvidersCreated) > 0 {
if err := tx.Where("id IN ?", backup.ProvidersCreated).Delete(&StorageProvider{}).Error; err != nil {
tx.Rollback()
return fmt.Errorf("failed to delete created providers: %v", err)
}
log.Printf("Deleted %d provider records created during migration", len(backup.ProvidersCreated))
}
// Then restore original config records
for _, config := range backup.Configs {
if err := tx.Save(&config).Error; err != nil {
tx.Rollback()
return fmt.Errorf("failed to restore config ID %d: %v", config.ID, err)
}
}
log.Printf("Restored %d transfer config records", len(backup.Configs))
// Commit the transaction
if err := tx.Commit().Error; err != nil {
return fmt.Errorf("failed to commit rollback transaction: %v", err)
}
log.Println("Rollback completed successfully")
return nil
}
// FormatMigrationReport generates a human-readable report of the migration results
func FormatMigrationReport(stats *MigrationStats) string {
if stats == nil {
return "No migration statistics available"
}
duration := stats.EndTime.Sub(stats.StartTime)
report := strings.Builder{}
report.WriteString("=== Provider Data Migration Report ===\n\n")
report.WriteString(fmt.Sprintf("Started: %s\n", stats.StartTime.Format(time.RFC3339)))
report.WriteString(fmt.Sprintf("Completed: %s\n", stats.EndTime.Format(time.RFC3339)))
report.WriteString(fmt.Sprintf("Duration: %s\n", duration))
report.WriteString(fmt.Sprintf("Total Configs: %d\n", stats.TotalConfigs))
report.WriteString(fmt.Sprintf("Source Providers: %d\n", stats.UniqueSourceProviders))
report.WriteString(fmt.Sprintf("Destination Providers: %d\n", stats.UniqueDestinationProviders))
report.WriteString(fmt.Sprintf("Providers Created: %d\n", stats.NewProvidersCreated))
report.WriteString(fmt.Sprintf("Configs Updated: %d\n", stats.ConfigsUpdated))
if len(stats.Errors) > 0 {
report.WriteString("\nErrors/Warnings:\n")
for i, err := range stats.Errors {
report.WriteString(fmt.Sprintf("%d. %s\n", i+1, err))
}
} else {
report.WriteString("\nNo errors or warnings reported.\n")
}
report.WriteString("\n=== End of Report ===\n")
return report.String()
}
@@ -0,0 +1,66 @@
package migrations
import (
"github.com/go-gormigrate/gormigrate/v2"
"gorm.io/gorm"
)
// AddStorageProviders adds the storage_providers table
func AddStorageProviders() *gormigrate.Migration {
return &gormigrate.Migration{
ID: "014_add_storage_providers",
Migrate: func(tx *gorm.DB) error {
// Create the storage_providers table
if err := tx.Exec(`CREATE TABLE IF NOT EXISTS storage_providers (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name VARCHAR(255) NOT NULL,
type VARCHAR(50) NOT NULL,
host VARCHAR(255),
port INTEGER DEFAULT 22,
username VARCHAR(255),
encrypted_password TEXT,
key_file TEXT,
bucket VARCHAR(255),
region VARCHAR(255),
access_key VARCHAR(255),
encrypted_secret_key TEXT,
endpoint VARCHAR(255),
share VARCHAR(255),
domain VARCHAR(255),
passive_mode BOOLEAN DEFAULT TRUE,
client_id VARCHAR(255),
encrypted_client_secret TEXT,
encrypted_refresh_token TEXT,
drive_id VARCHAR(255),
team_drive VARCHAR(255),
read_only BOOLEAN DEFAULT FALSE,
start_year INTEGER,
include_archived BOOLEAN DEFAULT FALSE,
use_builtin_auth BOOLEAN DEFAULT TRUE,
authenticated BOOLEAN DEFAULT FALSE,
created_by INTEGER NOT NULL,
created_at DATETIME,
updated_at DATETIME,
FOREIGN KEY (created_by) REFERENCES users(id),
UNIQUE(name, created_by)
)`).Error; err != nil {
return err
}
// Create index on type for faster filtering
if err := tx.Exec(`CREATE INDEX idx_storage_providers_type ON storage_providers(type)`).Error; err != nil {
return err
}
return nil
},
Rollback: func(tx *gorm.DB) error {
// Drop the storage_providers table
if err := tx.Exec(`DROP TABLE IF EXISTS storage_providers`).Error; err != nil {
return err
}
return nil
},
}
}
@@ -0,0 +1,55 @@
package migrations
import (
"github.com/go-gormigrate/gormigrate/v2"
"gorm.io/gorm"
)
// AddProviderRefsToTransferConfig adds the storage provider reference fields to the transfer_configs table
func AddProviderRefsToTransferConfig() *gormigrate.Migration {
return &gormigrate.Migration{
ID: "015_add_provider_refs_to_transfer_config",
Migrate: func(tx *gorm.DB) error {
// Add source_provider_id and destination_provider_id columns to transfer_configs table
if err := tx.Exec(`ALTER TABLE transfer_configs ADD COLUMN source_provider_id INTEGER REFERENCES storage_providers(id)`).Error; err != nil {
return err
}
if err := tx.Exec(`ALTER TABLE transfer_configs ADD COLUMN destination_provider_id INTEGER REFERENCES storage_providers(id)`).Error; err != nil {
return err
}
// Create indexes for better performance when joining with the storage_providers table
if err := tx.Exec(`CREATE INDEX idx_transfer_configs_source_provider_id ON transfer_configs(source_provider_id)`).Error; err != nil {
return err
}
if err := tx.Exec(`CREATE INDEX idx_transfer_configs_destination_provider_id ON transfer_configs(destination_provider_id)`).Error; err != nil {
return err
}
return nil
},
Rollback: func(tx *gorm.DB) error {
// Drop indexes first
if err := tx.Exec(`DROP INDEX IF EXISTS idx_transfer_configs_source_provider_id`).Error; err != nil {
return err
}
if err := tx.Exec(`DROP INDEX IF EXISTS idx_transfer_configs_destination_provider_id`).Error; err != nil {
return err
}
// Remove columns
if err := tx.Exec(`ALTER TABLE transfer_configs DROP COLUMN source_provider_id`).Error; err != nil {
return err
}
if err := tx.Exec(`ALTER TABLE transfer_configs DROP COLUMN destination_provider_id`).Error; err != nil {
return err
}
return nil
},
}
}
+2
View File
@@ -27,6 +27,8 @@ func GetMigrations(db *gorm.DB) *gormigrate.Gormigrate {
RecoverNotificationServicesRename(), // 012b
RecoverAuthProvidersRename(), // 012c
CleanupInvalidBooleans(), // 013
AddStorageProviders(), // 014
AddProviderRefsToTransferConfig(), // 015
)
return gormigrate.New(db, gormigrate.DefaultOptions, migrations)
+205
View File
@@ -0,0 +1,205 @@
package db
import (
"time"
)
// StorageProviderType defines the type of storage provider
type StorageProviderType string
const (
// Storage provider types
ProviderTypeSFTP StorageProviderType = "sftp"
ProviderTypeS3 StorageProviderType = "s3"
ProviderTypeOneDrive StorageProviderType = "onedrive"
ProviderTypeGoogleDrive StorageProviderType = "google_drive"
ProviderTypeGooglePhoto StorageProviderType = "google_photo"
ProviderTypeFTP StorageProviderType = "ftp"
ProviderTypeSMB StorageProviderType = "smb"
ProviderTypeHetzner StorageProviderType = "hetzner"
ProviderTypeLocal StorageProviderType = "local"
)
// StorageProvider represents a connection to a storage service
type StorageProvider struct {
ID uint `gorm:"primarykey" json:"id"`
Name string `gorm:"not null;uniqueIndex:idx_storage_providers_name_created_by" json:"name" form:"name"`
Type StorageProviderType `gorm:"not null" json:"type" form:"type"`
// Common fields
Host string `json:"host" form:"host"` // For server-based providers (SFTP, FTP, SMB)
Port int `gorm:"default:22" json:"port" form:"port"` // For server-based providers
Username string `json:"username" form:"username"` // Or AccessKey for S3
// Password is not stored in the database, only used for form input
Password string `gorm:"-" json:"-" form:"password"`
// These fields will be encrypted before storage
EncryptedPassword string `json:"-"` // Encrypted version of Password
KeyFile string `json:"key_file" form:"key_file"`
// S3 specific fields
Bucket string `json:"bucket" form:"bucket"`
Region string `json:"region" form:"region"`
AccessKey string `json:"access_key" form:"access_key"` // Alternative to Username for S3
// SecretKey is not stored in the database, only used for form input
SecretKey string `gorm:"-" json:"-" form:"secret_key"`
// Encrypted version of SecretKey
EncryptedSecretKey string `json:"-"`
Endpoint string `json:"endpoint" form:"endpoint"`
// SMB specific fields
Share string `json:"share" form:"share"`
Domain string `json:"domain" form:"domain"`
// FTP specific fields
PassiveMode *bool `gorm:"default:true" json:"passive_mode" form:"passive_mode"`
// OAuth-related fields for cloud providers (OneDrive, GoogleDrive, GooglePhoto)
ClientID string `json:"client_id" form:"client_id"`
// ClientSecret is not stored in the database, only used for form input
ClientSecret string `gorm:"-" json:"-" form:"client_secret"`
// Encrypted version of ClientSecret
EncryptedClientSecret string `json:"-"`
// RefreshToken is not stored in the database, only used for form input
RefreshToken string `gorm:"-" json:"-" form:"refresh_token"`
// Encrypted version of RefreshToken
EncryptedRefreshToken string `json:"-"`
// OAuth specific fields
DriveID string `json:"drive_id" form:"drive_id"` // For OneDrive
TeamDrive string `json:"team_drive" form:"team_drive"` // For Google Drive
ReadOnly *bool `json:"read_only" form:"read_only"` // For Google Photos
StartYear int `json:"start_year" form:"start_year"` // For Google Photos
IncludeArchived *bool `json:"include_archived" form:"include_archived"` // For Google Photos
// Security fields
UseBuiltinAuth *bool `gorm:"default:true" json:"use_builtin_auth" form:"use_builtin_auth"` // For OAuth services
// Status fields
Authenticated *bool `json:"authenticated"` // Whether auth is completed (for OAuth providers)
// Ownership and timestamps
CreatedBy uint `gorm:"not null;uniqueIndex:idx_storage_providers_name_created_by" json:"created_by"`
User User `gorm:"foreignkey:CreatedBy" json:"-"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// --- StorageProvider Helper Methods ---
// GetPassiveMode returns the value of PassiveMode with a default if nil
func (sp *StorageProvider) GetPassiveMode() bool {
if sp.PassiveMode == nil {
return true // Default to true if not set
}
return *sp.PassiveMode
}
// SetPassiveMode sets the PassiveMode field
func (sp *StorageProvider) SetPassiveMode(value bool) {
sp.PassiveMode = &value
}
// GetReadOnly returns the value of ReadOnly with a default if nil
func (sp *StorageProvider) GetReadOnly() bool {
if sp.ReadOnly == nil {
return false // Default to false if not set
}
return *sp.ReadOnly
}
// SetReadOnly sets the ReadOnly field
func (sp *StorageProvider) SetReadOnly(value bool) {
sp.ReadOnly = &value
}
// GetIncludeArchived returns the value of IncludeArchived with a default if nil
func (sp *StorageProvider) GetIncludeArchived() bool {
if sp.IncludeArchived == nil {
return false // Default to false if not set
}
return *sp.IncludeArchived
}
// SetIncludeArchived sets the IncludeArchived field
func (sp *StorageProvider) SetIncludeArchived(value bool) {
sp.IncludeArchived = &value
}
// GetUseBuiltinAuth returns the value of UseBuiltinAuth with a default if nil
func (sp *StorageProvider) GetUseBuiltinAuth() bool {
if sp.UseBuiltinAuth == nil {
return true // Default to true if not set
}
return *sp.UseBuiltinAuth
}
// SetUseBuiltinAuth sets the UseBuiltinAuth field
func (sp *StorageProvider) SetUseBuiltinAuth(value bool) {
sp.UseBuiltinAuth = &value
}
// GetAuthenticated returns the value of Authenticated with a default if nil
func (sp *StorageProvider) GetAuthenticated() bool {
if sp.Authenticated == nil {
return false // Default to false if not set
}
return *sp.Authenticated
}
// SetAuthenticated sets the Authenticated field
func (sp *StorageProvider) SetAuthenticated(value bool) {
sp.Authenticated = &value
}
// IsOAuthProvider returns true if the provider type requires OAuth authentication
func (sp *StorageProvider) IsOAuthProvider() bool {
return sp.Type == ProviderTypeOneDrive ||
sp.Type == ProviderTypeGoogleDrive ||
sp.Type == ProviderTypeGooglePhoto
}
// RequiresEncryption returns true if the provider has sensitive fields that need encryption
func (sp *StorageProvider) RequiresEncryption() bool {
// All provider types have some form of sensitive authentication that needs encryption
return true
}
// GetSensitiveFields returns a map of field names to values that need encryption
func (sp *StorageProvider) GetSensitiveFields() map[string]string {
sensitiveFields := make(map[string]string)
// Add fields based on provider type
switch sp.Type {
case ProviderTypeSFTP, ProviderTypeFTP, ProviderTypeSMB, ProviderTypeHetzner:
if sp.Password != "" {
sensitiveFields["Password"] = sp.Password
}
case ProviderTypeS3:
if sp.SecretKey != "" {
sensitiveFields["SecretKey"] = sp.SecretKey
}
case ProviderTypeOneDrive, ProviderTypeGoogleDrive, ProviderTypeGooglePhoto:
if sp.ClientSecret != "" {
sensitiveFields["ClientSecret"] = sp.ClientSecret
}
if sp.RefreshToken != "" {
sensitiveFields["RefreshToken"] = sp.RefreshToken
}
}
return sensitiveFields
}
// GetEncryptedFieldName returns the corresponding encrypted field name for a given sensitive field
func (sp *StorageProvider) GetEncryptedFieldName(fieldName string) string {
return "Encrypted" + fieldName
}
+42
View File
@@ -0,0 +1,42 @@
package db
import (
"time"
)
// ConnectorError represents different types of connection errors
type ConnectorError struct {
Code string
Message string
Err error
}
// ConnectionResult contains the result of a connection test
type ConnectionResult struct {
Success bool
Message string
Error *ConnectorError
Timestamp time.Time
}
// Common error codes
const (
ErrorCodeUnknown = "unknown"
ErrorCodeTimeout = "timeout"
ErrorCodeAuthentication = "authentication"
ErrorCodeConnection = "connection"
ErrorCodeResourceNotFound = "resource_not_found"
ErrorCodeInvalidParams = "invalid_params"
ErrorCodePermission = "permission"
ErrorCodeNetwork = "network"
)
// Error returns the error message
func (e *ConnectorError) Error() string {
return e.Message
}
// Unwrap returns the underlying error
func (e *ConnectorError) Unwrap() error {
return e.Err
}
@@ -0,0 +1,416 @@
package db
import (
"testing"
"time"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
var testDB *gorm.DB
// setupTestDB sets up a SQLite in-memory database for testing
func setupTestDB(t *testing.T) *DB {
var err error
testDB, err = gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
if err != nil {
t.Fatalf("Failed to open in-memory SQLite database: %v", err)
}
// Create a minimal TransferConfig struct for testing
type TransferConfig struct {
ID uint `gorm:"primarykey"`
SourceProviderID uint `gorm:"index"`
DestinationProviderID uint `gorm:"index"`
}
// Create the necessary tables
err = testDB.AutoMigrate(&StorageProvider{}, &TransferConfig{})
if err != nil {
t.Fatalf("Failed to migrate tables: %v", err)
}
return &DB{DB: testDB}
}
// cleanupTestDB cleans up the test database after each test
func cleanupTestDB(t *testing.T) {
sqlDB, err := testDB.DB()
if err != nil {
t.Fatalf("Failed to get SQL DB: %v", err)
}
sqlDB.Close()
}
// TestStorageProviderCRUD tests the complete CRUD cycle for a StorageProvider
func TestStorageProviderCRUD(t *testing.T) {
db := setupTestDB(t)
defer cleanupTestDB(t)
// Create a test provider
provider := &StorageProvider{
Name: "Test SFTP",
Type: ProviderTypeSFTP,
Host: "example.com",
Port: 22,
Username: "user",
Password: "password",
CreatedBy: 1,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Test Create
err := db.CreateStorageProvider(provider)
if err != nil {
t.Fatalf("Failed to create storage provider: %v", err)
}
if provider.ID == 0 {
t.Fatal("Expected provider ID to be set after creation")
}
// Test Get
retrievedProvider, err := db.GetStorageProvider(provider.ID)
if err != nil {
t.Fatalf("Failed to get storage provider: %v", err)
}
if retrievedProvider.ID != provider.ID {
t.Errorf("Expected provider ID %d, got %d", provider.ID, retrievedProvider.ID)
}
if retrievedProvider.Name != "Test SFTP" {
t.Errorf("Expected name 'Test SFTP', got '%s'", retrievedProvider.Name)
}
if retrievedProvider.Type != ProviderTypeSFTP {
t.Errorf("Expected type '%s', got '%s'", ProviderTypeSFTP, retrievedProvider.Type)
}
// Test Update
retrievedProvider.Name = "Updated SFTP"
retrievedProvider.Host = "updated.example.com"
// Make sure we keep the required fields for validation
retrievedProvider.Port = 22
retrievedProvider.Username = "user"
retrievedProvider.Password = "password"
err = db.UpdateStorageProvider(retrievedProvider)
if err != nil {
t.Fatalf("Failed to update storage provider: %v", err)
}
// Verify update
updatedProvider, err := db.GetStorageProvider(provider.ID)
if err != nil {
t.Fatalf("Failed to get updated storage provider: %v", err)
}
if updatedProvider.Name != "Updated SFTP" {
t.Errorf("Expected updated name 'Updated SFTP', got '%s'", updatedProvider.Name)
}
if updatedProvider.Host != "updated.example.com" {
t.Errorf("Expected updated host 'updated.example.com', got '%s'", updatedProvider.Host)
}
// Test Delete
err = db.DeleteStorageProvider(provider.ID)
if err != nil {
t.Fatalf("Failed to delete storage provider: %v", err)
}
// Verify deletion
_, err = db.GetStorageProvider(provider.ID)
if err == nil {
t.Error("Expected error when getting deleted provider, got nil")
}
}
// TestStorageProviderGetAll tests retrieving all storage providers for a user
func TestStorageProviderGetAll(t *testing.T) {
db := setupTestDB(t)
defer cleanupTestDB(t)
// Create multiple providers for the same user
providers := []*StorageProvider{
{
Name: "SFTP Provider",
Type: ProviderTypeSFTP,
Host: "sftp.example.com",
Port: 22,
Username: "sftpuser",
Password: "pass",
CreatedBy: 1,
},
{
Name: "S3 Provider",
Type: ProviderTypeS3,
AccessKey: "accesskey",
SecretKey: "secretkey",
Region: "us-west-1",
CreatedBy: 1,
},
{
Name: "OneDrive Provider",
Type: ProviderTypeOneDrive,
ClientID: "clientid",
ClientSecret: "clientsecret",
CreatedBy: 1,
},
{
Name: "Another User's Provider",
Type: ProviderTypeSFTP,
Host: "other.example.com",
Port: 22, // Added required port for SFTP
Username: "otheruser", // Added required username for SFTP
Password: "otherpass", // Added required password for SFTP
CreatedBy: 2, // Different user
},
}
// Create all providers
for _, p := range providers {
err := db.CreateStorageProvider(p)
if err != nil {
t.Fatalf("Failed to create provider %s: %v", p.Name, err)
}
}
// Test GetStorageProviders
userProviders, err := db.GetStorageProviders(1)
if err != nil {
t.Fatalf("Failed to get storage providers: %v", err)
}
// Check the results
if len(userProviders) != 3 {
t.Errorf("Expected 3 providers for user 1, got %d", len(userProviders))
}
}
// TestStorageProviderGetByType tests retrieving storage providers by type
func TestStorageProviderGetByType(t *testing.T) {
db := setupTestDB(t)
defer cleanupTestDB(t)
// Create providers of different types
providers := []*StorageProvider{
{
Name: "SFTP Provider 1",
Type: ProviderTypeSFTP,
Host: "sftp1.example.com",
Port: 22, // Added required port for SFTP
Username: "user1", // Added required username for SFTP
Password: "pass1", // Added required password for SFTP
CreatedBy: 1,
},
{
Name: "SFTP Provider 2",
Type: ProviderTypeSFTP,
Host: "sftp2.example.com",
Port: 22, // Added required port for SFTP
Username: "user2", // Added required username for SFTP
Password: "pass2", // Added required password for SFTP
CreatedBy: 1,
},
{
Name: "S3 Provider",
Type: ProviderTypeS3,
AccessKey: "accesskey",
SecretKey: "secretkey", // Added required secret key for S3
Region: "us-west-1", // Added required region for S3
CreatedBy: 1,
},
}
// Create all providers
for _, p := range providers {
err := db.CreateStorageProvider(p)
if err != nil {
t.Fatalf("Failed to create provider %s: %v", p.Name, err)
}
}
// Test GetStorageProvidersByType
sftpProviders, err := db.GetStorageProvidersByType(1, ProviderTypeSFTP)
if err != nil {
t.Fatalf("Failed to get SFTP providers: %v", err)
}
// Check the results
if len(sftpProviders) != 2 {
t.Errorf("Expected 2 SFTP providers, got %d", len(sftpProviders))
}
for _, p := range sftpProviders {
if p.Type != ProviderTypeSFTP {
t.Errorf("Expected provider type SFTP, got %s", p.Type)
}
}
}
// TestStorageProviderValidationOnSave tests that validation is called before saving
func TestStorageProviderValidationOnSave(t *testing.T) {
db := setupTestDB(t)
defer cleanupTestDB(t)
// Create a provider with invalid data (missing host for SFTP)
invalidProvider := &StorageProvider{
Name: "Invalid SFTP",
Type: ProviderTypeSFTP,
Port: 22,
Username: "user",
Password: "pass",
CreatedBy: 1,
}
// Test CreateStorageProvider with validation
err := db.CreateStorageProvider(invalidProvider)
if err == nil {
t.Fatal("Expected validation error for invalid provider, got nil")
}
}
// TestStorageProviderCount tests counting providers for a user
func TestStorageProviderCount(t *testing.T) {
db := setupTestDB(t)
defer cleanupTestDB(t)
// Create multiple providers for different users
providers := []*StorageProvider{
{
Name: "User 1 Provider 1",
Type: ProviderTypeSFTP,
Host: "host1.example.com",
Port: 22, // Added required port for SFTP
Username: "user1", // Added required username for SFTP
Password: "pass1", // Added required password for SFTP
CreatedBy: 1,
},
{
Name: "User 1 Provider 2",
Type: ProviderTypeS3,
AccessKey: "accesskey",
SecretKey: "secretkey", // Added required secret key for S3
Region: "us-west-1", // Added required region for S3
CreatedBy: 1,
},
{
Name: "User 2 Provider",
Type: ProviderTypeSFTP,
Host: "host2.example.com",
Port: 22, // Added required port for SFTP
Username: "user2", // Added required username for SFTP
Password: "pass2", // Added required password for SFTP
CreatedBy: 2,
},
}
// Create all providers
for _, p := range providers {
// Skip validation for this test since we're just testing count
err := testDB.Create(p).Error
if err != nil {
t.Fatalf("Failed to create provider %s: %v", p.Name, err)
}
}
// Test CountStorageProviders
count, err := db.CountStorageProviders(1)
if err != nil {
t.Fatalf("Failed to count storage providers: %v", err)
}
// Check the result
if count != 2 {
t.Errorf("Expected count 2 for user 1, got %d", count)
}
count, err = db.CountStorageProviders(2)
if err != nil {
t.Fatalf("Failed to count storage providers: %v", err)
}
// Check the result
if count != 1 {
t.Errorf("Expected count 1 for user 2, got %d", count)
}
}
// TestHelperMethods tests the helper methods on StorageProvider
func TestHelperMethods(t *testing.T) {
// Test GetPassiveMode and SetPassiveMode
t.Run("PassiveMode", func(t *testing.T) {
provider := &StorageProvider{}
// Default value
if !provider.GetPassiveMode() {
t.Error("Expected default PassiveMode to be true")
}
// Set to false
provider.SetPassiveMode(false)
if provider.GetPassiveMode() {
t.Error("Expected PassiveMode to be false after setting")
}
// Set to true
provider.SetPassiveMode(true)
if !provider.GetPassiveMode() {
t.Error("Expected PassiveMode to be true after setting")
}
})
// Test GetReadOnly and SetReadOnly
t.Run("ReadOnly", func(t *testing.T) {
provider := &StorageProvider{}
// Default value
if provider.GetReadOnly() {
t.Error("Expected default ReadOnly to be false")
}
// Set to true
provider.SetReadOnly(true)
if !provider.GetReadOnly() {
t.Error("Expected ReadOnly to be true after setting")
}
})
// Test GetAuthenticated and SetAuthenticated
t.Run("Authenticated", func(t *testing.T) {
provider := &StorageProvider{}
// Default value
if provider.GetAuthenticated() {
t.Error("Expected default Authenticated to be false")
}
// Set to true
provider.SetAuthenticated(true)
if !provider.GetAuthenticated() {
t.Error("Expected Authenticated to be true after setting")
}
})
}
// TestIsOAuthProvider tests the IsOAuthProvider method
func TestIsOAuthProvider(t *testing.T) {
tests := []struct {
providerType StorageProviderType
isOAuth bool
}{
{ProviderTypeSFTP, false},
{ProviderTypeS3, false},
{ProviderTypeFTP, false},
{ProviderTypeSMB, false},
{ProviderTypeOneDrive, true},
{ProviderTypeGoogleDrive, true},
{ProviderTypeGooglePhoto, true},
{ProviderTypeLocal, false},
}
for _, tt := range tests {
t.Run(string(tt.providerType), func(t *testing.T) {
provider := &StorageProvider{Type: tt.providerType}
if provider.IsOAuthProvider() != tt.isOAuth {
t.Errorf("IsOAuthProvider() for %s = %v, want %v", tt.providerType, provider.IsOAuthProvider(), tt.isOAuth)
}
})
}
}
+83
View File
@@ -0,0 +1,83 @@
package db
import (
"fmt"
)
// --- StorageProvider Store Methods ---
// CreateStorageProvider creates a new storage provider record
func (db *DB) CreateStorageProvider(provider *StorageProvider) error {
return db.Create(provider).Error
}
// GetStorageProviders retrieves all storage providers for a user
func (db *DB) GetStorageProviders(userID uint) ([]StorageProvider, error) {
var providers []StorageProvider
err := db.Where("created_by = ?", userID).Find(&providers).Error
return providers, err
}
// GetStorageProvidersByType retrieves all storage providers of a specific type for a user
func (db *DB) GetStorageProvidersByType(userID uint, providerType StorageProviderType) ([]StorageProvider, error) {
var providers []StorageProvider
err := db.Where("created_by = ? AND type = ?", userID, providerType).Find(&providers).Error
return providers, err
}
// GetStorageProvider retrieves a single storage provider by ID
func (db *DB) GetStorageProvider(id uint) (*StorageProvider, error) {
var provider StorageProvider
err := db.First(&provider, id).Error
if err != nil {
return nil, err
}
return &provider, nil
}
// GetStorageProviderType retrieves the type of a storage provider by ID
func (db *DB) GetStorageProviderType(id uint) (StorageProviderType, error) {
var provider StorageProvider
err := db.First(&provider, id).Error
return provider.Type, err
}
// GetStorageProviderWithOwnerCheck retrieves a single storage provider by ID with owner check
func (db *DB) GetStorageProviderWithOwnerCheck(id uint, userID uint) (*StorageProvider, error) {
var provider StorageProvider
err := db.Where("id = ? AND created_by = ?", id, userID).First(&provider).Error
if err != nil {
return nil, err
}
return &provider, nil
}
// UpdateStorageProvider updates an existing storage provider record
func (db *DB) UpdateStorageProvider(provider *StorageProvider) error {
return db.Save(provider).Error
}
// DeleteStorageProvider deletes a storage provider record after checking dependencies
func (db *DB) DeleteStorageProvider(id uint) error {
// First check if any transfer configs are using this provider
var count int64
if err := db.Model(&TransferConfig{}).
Where("source_provider_id = ? OR destination_provider_id = ?", id, id).
Count(&count).Error; err != nil {
return fmt.Errorf("failed to check for dependent transfer configs: %v", err)
}
if count > 0 {
return fmt.Errorf("cannot delete provider: %d transfer configurations are using this provider", count)
}
// Delete the provider
return db.Delete(&StorageProvider{}, id).Error
}
// CountStorageProviders counts the number of storage providers for a user
func (db *DB) CountStorageProviders(userID uint) (int64, error) {
var count int64
err := db.Model(&StorageProvider{}).Where("created_by = ?", userID).Count(&count).Error
return count, err
}
+241
View File
@@ -0,0 +1,241 @@
package db
import (
"errors"
"fmt"
"log"
"strings"
"gorm.io/gorm"
)
// ValidateStorageProvider validates a storage provider based on its type
func (sp *StorageProvider) Validate() error {
// Special case - if we have an empty struct or just an ID (could happen during GORM operations like foreign key checks)
if (sp.ID > 0 && sp.Name == "" && sp.Type == "") || (sp.ID == 0 && sp.Name == "" && sp.Type == "") {
log.Printf("Skipping validation for StorageProvider: ID=%d without other data (likely a reference check)", sp.ID)
return nil
}
// Common validations
if strings.TrimSpace(sp.Name) == "" {
return errors.New("provider name cannot be empty")
}
// Type-specific validations
switch sp.Type {
case ProviderTypeSFTP, ProviderTypeHetzner:
return sp.validateSFTP()
case ProviderTypeS3:
return sp.validateS3()
case ProviderTypeFTP:
return sp.validateFTP()
case ProviderTypeSMB:
return sp.validateSMB()
case ProviderTypeOneDrive:
return sp.validateOneDrive()
case ProviderTypeGoogleDrive:
return sp.validateGoogleDrive()
case ProviderTypeGooglePhoto:
return sp.validateGooglePhoto()
case ProviderTypeLocal:
return sp.validateLocal()
default:
return fmt.Errorf("unsupported provider type: %s", sp.Type)
}
}
// validateSFTP validates SFTP-specific fields
func (sp *StorageProvider) validateSFTP() error {
if strings.TrimSpace(sp.Host) == "" {
return errors.New("host is required for SFTP provider")
}
if sp.Port <= 0 {
return errors.New("invalid port for SFTP provider")
}
if strings.TrimSpace(sp.Username) == "" {
return errors.New("username is required for SFTP provider")
}
// Either password or key file must be provided
if strings.TrimSpace(sp.Password) == "" && strings.TrimSpace(sp.EncryptedPassword) == "" && strings.TrimSpace(sp.KeyFile) == "" {
return errors.New("either password or key file is required for SFTP provider")
}
return nil
}
// validateS3 validates S3-specific fields
func (sp *StorageProvider) validateS3() error {
// For S3, either AccessKey or Username is used
if strings.TrimSpace(sp.AccessKey) == "" && strings.TrimSpace(sp.Username) == "" {
return errors.New("access key is required for S3 provider")
}
// Either SecretKey or EncryptedSecretKey must be provided
if strings.TrimSpace(sp.SecretKey) == "" && strings.TrimSpace(sp.EncryptedSecretKey) == "" {
return errors.New("secret key is required for S3 provider")
}
// Region is required for most S3 providers
if strings.TrimSpace(sp.Region) == "" {
return errors.New("region is required for S3 provider")
}
return nil
}
// validateFTP validates FTP-specific fields
func (sp *StorageProvider) validateFTP() error {
if strings.TrimSpace(sp.Host) == "" {
return errors.New("host is required for FTP provider")
}
if sp.Port <= 0 {
return errors.New("invalid port for FTP provider")
}
if strings.TrimSpace(sp.Username) == "" {
return errors.New("username is required for FTP provider")
}
// Either password or encrypted password must be provided
if strings.TrimSpace(sp.Password) == "" && strings.TrimSpace(sp.EncryptedPassword) == "" {
return errors.New("password is required for FTP provider")
}
return nil
}
// validateSMB validates SMB-specific fields
func (sp *StorageProvider) validateSMB() error {
if strings.TrimSpace(sp.Host) == "" {
return errors.New("host is required for SMB provider")
}
if strings.TrimSpace(sp.Share) == "" {
return errors.New("share is required for SMB provider")
}
if strings.TrimSpace(sp.Username) == "" {
return errors.New("username is required for SMB provider")
}
// Either password or encrypted password must be provided
if strings.TrimSpace(sp.Password) == "" && strings.TrimSpace(sp.EncryptedPassword) == "" {
return errors.New("password is required for SMB provider")
}
return nil
}
// validateOneDrive validates OneDrive-specific fields
func (sp *StorageProvider) validateOneDrive() error {
if strings.TrimSpace(sp.ClientID) == "" {
return errors.New("client ID is required for OneDrive provider")
}
// Either ClientSecret or EncryptedClientSecret must be provided
if strings.TrimSpace(sp.ClientSecret) == "" && strings.TrimSpace(sp.EncryptedClientSecret) == "" {
return errors.New("client secret is required for OneDrive provider")
}
// For authenticated providers, RefreshToken must be set
if sp.GetAuthenticated() && strings.TrimSpace(sp.EncryptedRefreshToken) == "" && strings.TrimSpace(sp.RefreshToken) == "" {
return errors.New("refresh token is required for authenticated OneDrive provider")
}
return nil
}
// validateGoogleDrive validates Google Drive-specific fields
func (sp *StorageProvider) validateGoogleDrive() error {
// If not using builtin auth, ClientID and ClientSecret are required
if !sp.GetUseBuiltinAuth() {
if strings.TrimSpace(sp.ClientID) == "" {
return errors.New("client ID is required for Google Drive provider when not using builtin auth")
}
// Either ClientSecret or EncryptedClientSecret must be provided
if strings.TrimSpace(sp.ClientSecret) == "" && strings.TrimSpace(sp.EncryptedClientSecret) == "" {
return errors.New("client secret is required for Google Drive provider when not using builtin auth")
}
}
// For authenticated providers, RefreshToken must be set
if sp.GetAuthenticated() && strings.TrimSpace(sp.EncryptedRefreshToken) == "" && strings.TrimSpace(sp.RefreshToken) == "" {
return errors.New("refresh token is required for authenticated Google Drive provider")
}
return nil
}
// validateGooglePhoto validates Google Photos-specific fields
func (sp *StorageProvider) validateGooglePhoto() error {
// Similar to Google Drive
return sp.validateGoogleDrive()
}
// validateLocal validates Local-specific fields
func (sp *StorageProvider) validateLocal() error {
// Local providers don't need additional validation
return nil
}
// BeforeSave is a GORM hook that runs before saving the provider
func (sp *StorageProvider) BeforeSave(tx *gorm.DB) error {
// Check if this is a reference check by examining the GORM operation
if tx.Statement.SQL.String() == "" {
// No explicit SQL means this might be part of a preload or association check
// Case 1: Empty struct (as you already have)
if sp.ID == 0 && sp.Name == "" && sp.Type == "" {
log.Printf("BeforeSave: Skipping validation for empty StorageProvider")
return nil
}
// Case 2: ID-only struct (foreign key reference check)
if sp.ID > 0 && sp.Name == "" && sp.Type == "" {
log.Printf("BeforeSave: Skipping validation for StorageProvider ID=%d (reference check)", sp.ID)
return nil
}
// Case 3: Minimal data loaded from database for relationship check
// Check if only a few fields are populated (typically ID and maybe a couple others)
populatedFields := 0
if sp.ID > 0 {
populatedFields++
}
if sp.Name != "" {
populatedFields++
}
if string(sp.Type) != "" {
populatedFields++
}
if sp.Host != "" {
populatedFields++
}
if sp.Username != "" {
populatedFields++
}
// If we have just a few populated fields, it's likely a reference check
if populatedFields <= 3 {
log.Printf("BeforeSave: Skipping validation for partially loaded StorageProvider ID=%d (likely reference check)", sp.ID)
return nil
}
}
// Check if this is called from a foreign key operation on another model
stmt := tx.Statement
if stmt.Schema != nil && stmt.Schema.Table != "storage_providers" {
log.Printf("BeforeSave: Skipping validation for StorageProvider ID=%d (called from %s table operation)",
sp.ID, stmt.Schema.Table)
return nil
}
log.Printf("BeforeSave: Validating StorageProvider: ID=%d, Name=%s, Type=%s", sp.ID, sp.Name, sp.Type)
return sp.Validate()
}
@@ -0,0 +1,141 @@
package db
import (
"testing"
)
// TestStorageProviderValidation tests the validation of storage providers
func TestStorageProviderValidation(t *testing.T) {
tests := []struct {
name string
provider StorageProvider
wantError bool
}{
{
name: "Valid SFTP provider",
provider: StorageProvider{
Name: "Test SFTP",
Type: ProviderTypeSFTP,
Host: "example.com",
Port: 22,
Username: "user",
Password: "pass",
},
wantError: false,
},
{
name: "Invalid SFTP provider - missing host",
provider: StorageProvider{
Name: "Test SFTP",
Type: ProviderTypeSFTP,
Port: 22,
Username: "user",
Password: "pass",
},
wantError: true,
},
{
name: "Valid S3 provider",
provider: StorageProvider{
Name: "Test S3",
Type: ProviderTypeS3,
AccessKey: "accesskey",
SecretKey: "secretkey",
Region: "us-west-1",
},
wantError: false,
},
{
name: "Valid OneDrive provider",
provider: StorageProvider{
Name: "Test OneDrive",
Type: ProviderTypeOneDrive,
ClientID: "clientid",
ClientSecret: "clientsecret",
},
wantError: false,
},
// Testing all provider types to ensure they're correctly recognized in the switch statement
{
name: "Valid Hetzner provider",
provider: StorageProvider{
Name: "Test Hetzner",
Type: ProviderTypeHetzner,
Host: "example.com",
Port: 22,
Username: "user",
Password: "pass",
},
wantError: false,
},
{
name: "Valid FTP provider",
provider: StorageProvider{
Name: "Test FTP",
Type: ProviderTypeFTP,
Host: "example.com",
Port: 21,
Username: "user",
Password: "pass",
},
wantError: false,
},
{
name: "Valid SMB provider",
provider: StorageProvider{
Name: "Test SMB",
Type: ProviderTypeSMB,
Host: "example.com",
Share: "share",
Username: "user",
Password: "pass",
},
wantError: false,
},
{
name: "Valid Google Drive provider",
provider: StorageProvider{
Name: "Test Google Drive",
Type: ProviderTypeGoogleDrive,
ClientID: "clientid",
ClientSecret: "clientsecret",
},
wantError: false,
},
{
name: "Valid Google Photo provider",
provider: StorageProvider{
Name: "Test Google Photo",
Type: ProviderTypeGooglePhoto,
ClientID: "clientid",
ClientSecret: "clientsecret",
},
wantError: false,
},
{
name: "Valid Local provider",
provider: StorageProvider{
Name: "Test Local",
Type: ProviderTypeLocal,
},
wantError: false,
},
{
name: "Invalid provider type",
provider: StorageProvider{
Name: "Test Invalid",
Type: "invalid_type",
},
wantError: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := tt.provider.Validate()
if (err != nil) != tt.wantError {
t.Errorf("Validate() error = %v, wantError %v", err, tt.wantError)
}
})
}
}
+544
View File
@@ -1,6 +1,7 @@
package db
import (
"fmt"
"time"
)
@@ -15,6 +16,9 @@ type TransferConfig struct {
SourceUser string `form:"source_user"`
SourcePassword string `form:"source_password" gorm:"-"` // Not stored in DB, only used for form
SourceKeyFile string `form:"source_key_file"`
// Source provider reference
SourceProviderID *uint `form:"source_provider_id"`
SourceProvider *StorageProvider `gorm:"foreignKey:SourceProviderID" json:"-"`
// S3 source fields
SourceBucket string `form:"source_bucket"`
SourceRegion string `form:"source_region"`
@@ -45,6 +49,9 @@ type TransferConfig struct {
DestUser string `form:"dest_user"`
DestPassword string `form:"dest_password" gorm:"-"` // Not stored in DB, only used for form
DestKeyFile string `form:"dest_key_file"`
// Destination provider reference
DestinationProviderID *uint `form:"destination_provider_id"`
DestinationProvider *StorageProvider `gorm:"foreignKey:DestinationProviderID" json:"-"`
// S3 destination fields
DestBucket string `form:"dest_bucket"`
DestRegion string `form:"dest_region"`
@@ -198,3 +205,540 @@ func (tc *TransferConfig) GetUseBuiltinAuthDest() bool {
func (tc *TransferConfig) SetUseBuiltinAuthDest(value bool) {
tc.UseBuiltinAuthDest = &value
}
// --- Provider Reference Methods ---
// IsUsingSourceProviderReference returns true if this config is using a source provider reference
func (tc *TransferConfig) IsUsingSourceProviderReference() bool {
return tc.SourceProviderID != nil && *tc.SourceProviderID > 0
}
// IsUsingDestinationProviderReference returns true if this config is using a destination provider reference
func (tc *TransferConfig) IsUsingDestinationProviderReference() bool {
return tc.DestinationProviderID != nil && *tc.DestinationProviderID > 0
}
// IsUsingProviderReferences returns true if this config is using provider references for both source and destination
func (tc *TransferConfig) IsUsingProviderReferences() bool {
return tc.IsUsingSourceProviderReference() && tc.IsUsingDestinationProviderReference()
}
// SetSourceProvider sets the source provider and ID fields
func (tc *TransferConfig) SetSourceProvider(provider *StorageProvider) {
if provider == nil || provider.ID == 0 {
tc.SourceProviderID = nil
tc.SourceProvider = nil
return
}
// Create a new uint pointer to avoid shared memory issues
newID := provider.ID
tc.SourceProviderID = &newID
tc.SourceProvider = provider
// Set the source type to match the provider type if not already set
if provider.Type != "" {
tc.SourceType = string(provider.Type)
}
}
// SetDestinationProvider sets the destination provider and ID fields
func (tc *TransferConfig) SetDestinationProvider(provider *StorageProvider) {
if provider == nil || provider.ID == 0 {
tc.DestinationProviderID = nil
tc.DestinationProvider = nil
return
}
// Create a new uint pointer to avoid shared memory issues
newID := provider.ID
tc.DestinationProviderID = &newID
tc.DestinationProvider = provider
// Set the destination type to match the provider type if not already set
if provider.Type != "" {
tc.DestinationType = string(provider.Type)
}
}
// EnsureProvidersLoaded ensures that both source and destination providers are loaded if references are used
func (tc *TransferConfig) EnsureProvidersLoaded(db interface{}) error {
if db == nil {
return fmt.Errorf("database interface is required to load providers")
}
// Try to load source provider if needed
if tc.IsUsingSourceProviderReference() && tc.SourceProvider == nil {
switch dbImpl := db.(type) {
case *DB:
provider, err := dbImpl.GetStorageProvider(*tc.SourceProviderID)
if err != nil {
return fmt.Errorf("failed to load source provider (ID %d): %w", *tc.SourceProviderID, err)
}
tc.SetSourceProvider(provider)
default:
return fmt.Errorf("invalid database interface for loading source provider")
}
}
// Try to load destination provider if needed
if tc.IsUsingDestinationProviderReference() && tc.DestinationProvider == nil {
switch dbImpl := db.(type) {
case *DB:
provider, err := dbImpl.GetStorageProvider(*tc.DestinationProviderID)
if err != nil {
return fmt.Errorf("failed to load destination provider (ID %d): %w", *tc.DestinationProviderID, err)
}
tc.SetDestinationProvider(provider)
default:
return fmt.Errorf("invalid database interface for loading destination provider")
}
}
return nil
}
// ValidateProviderConfiguration validates that the provider configuration is consistent
func (tc *TransferConfig) ValidateProviderConfiguration() error {
// Validate source provider configuration
if tc.IsUsingSourceProviderReference() {
if tc.SourceProvider == nil {
return fmt.Errorf("source provider reference set but provider is nil")
}
if tc.SourceProviderID == nil || *tc.SourceProviderID != tc.SourceProvider.ID {
return fmt.Errorf("source provider ID mismatch")
}
if tc.SourceType != string(tc.SourceProvider.Type) {
return fmt.Errorf("source type mismatch: config has %s but provider has %s", tc.SourceType, tc.SourceProvider.Type)
}
}
// Validate destination provider configuration
if tc.IsUsingDestinationProviderReference() {
if tc.DestinationProvider == nil {
return fmt.Errorf("destination provider reference set but provider is nil")
}
if tc.DestinationProviderID == nil || *tc.DestinationProviderID != tc.DestinationProvider.ID {
return fmt.Errorf("destination provider ID mismatch")
}
if tc.DestinationType != string(tc.DestinationProvider.Type) {
return fmt.Errorf("destination type mismatch: config has %s but provider has %s", tc.DestinationType, tc.DestinationProvider.Type)
}
}
return nil
}
// GetSourceCredentials returns credential information for the source, either directly or from the provider
// If db is provided, it will try to load the provider from the database if needed
func (tc *TransferConfig) GetSourceCredentials(db interface{}) (map[string]interface{}, error) {
creds := make(map[string]interface{})
fmt.Printf("DEBUG GetSourceCreds Start: ProviderID=%v, HasProvider=%v\n",
tc.SourceProviderID,
tc.SourceProvider != nil)
// If using provider reference and provider is loaded
if tc.IsUsingSourceProviderReference() {
// Try to load provider from database if we have a valid ID but no provider
if tc.SourceProvider == nil && db != nil {
// Try different types of DB interfaces to load the provider
switch dbImpl := db.(type) {
case *DB:
provider, err := dbImpl.GetStorageProvider(*tc.SourceProviderID)
if err != nil {
return nil, fmt.Errorf("failed to load source provider (ID %d): %w", *tc.SourceProviderID, err)
}
tc.SourceProvider = provider
case interface {
GetStorageProvider(id uint) (*StorageProvider, error)
}:
provider, err := dbImpl.GetStorageProvider(*tc.SourceProviderID)
if err != nil {
return nil, fmt.Errorf("failed to load source provider (ID %d): %w", *tc.SourceProviderID, err)
}
tc.SourceProvider = provider
default:
return nil, fmt.Errorf("source provider not loaded and db interface cannot load providers")
}
}
// If we still don't have a provider or it has no ID, return error
if tc.SourceProvider == nil || tc.SourceProvider.ID == 0 {
return nil, fmt.Errorf("failed to load valid source provider (ID %d)", *tc.SourceProviderID)
}
if tc.SourceProvider != nil {
fmt.Printf("DEBUG Provider Details:\n"+
" ID: %v\n"+
" Type: %v\n"+
" Host: %v\n"+
" Port: %v\n"+
" Username: %v\n"+
" HasEncryptedPassword: %v\n"+
" HasKeyFile: %v\n"+
" HasSecretKey: %v\n"+
" HasClientSecret: %v\n"+
" HasRefreshToken: %v\n",
tc.SourceProvider.ID,
tc.SourceProvider.Type,
tc.SourceProvider.Host,
tc.SourceProvider.Port,
tc.SourceProvider.Username,
tc.SourceProvider.EncryptedPassword != "",
tc.SourceProvider.KeyFile != "",
tc.SourceProvider.EncryptedSecretKey != "",
tc.SourceProvider.EncryptedClientSecret != "",
tc.SourceProvider.EncryptedRefreshToken != "")
}
// Copy credentials from provider
creds["type"] = tc.SourceProvider.Type
creds["host"] = tc.SourceProvider.Host
creds["port"] = tc.SourceProvider.Port
creds["username"] = tc.SourceProvider.Username
creds["encrypted_password"] = tc.SourceProvider.EncryptedPassword
creds["key_file"] = tc.SourceProvider.KeyFile
// Handle S3 fields
creds["bucket"] = tc.SourceProvider.Bucket
creds["region"] = tc.SourceProvider.Region
creds["access_key"] = tc.SourceProvider.AccessKey
creds["encrypted_secret_key"] = tc.SourceProvider.EncryptedSecretKey
creds["endpoint"] = tc.SourceProvider.Endpoint
// Handle SMB fields
creds["share"] = tc.SourceProvider.Share
creds["domain"] = tc.SourceProvider.Domain
// Handle FTP fields
if tc.SourceProvider.PassiveMode != nil {
creds["passive_mode"] = *tc.SourceProvider.PassiveMode
}
// Handle OAuth fields
creds["client_id"] = tc.SourceProvider.ClientID
creds["encrypted_client_secret"] = tc.SourceProvider.EncryptedClientSecret
creds["encrypted_refresh_token"] = tc.SourceProvider.EncryptedRefreshToken
creds["drive_id"] = tc.SourceProvider.DriveID
creds["team_drive"] = tc.SourceProvider.TeamDrive
if tc.SourceProvider.ReadOnly != nil {
creds["read_only"] = *tc.SourceProvider.ReadOnly
}
creds["start_year"] = tc.SourceProvider.StartYear
if tc.SourceProvider.IncludeArchived != nil {
creds["include_archived"] = *tc.SourceProvider.IncludeArchived
}
if tc.SourceProvider.UseBuiltinAuth != nil {
creds["use_builtin_auth"] = *tc.SourceProvider.UseBuiltinAuth
}
if tc.SourceProvider.Authenticated != nil {
creds["authenticated"] = *tc.SourceProvider.Authenticated
}
fmt.Printf("DEBUG Final Provider Creds:\n"+
" type: %v\n"+
" host: %v\n"+
" port: %v\n"+
" username: %v\n"+
" has_encrypted_password: %v\n"+
" has_key_file: %v\n"+
" has_encrypted_secret_key: %v\n"+
" has_encrypted_client_secret: %v\n",
creds["type"],
creds["host"],
creds["port"],
creds["username"],
creds["encrypted_password"] != "",
creds["key_file"] != "",
creds["encrypted_secret_key"] != "",
creds["encrypted_client_secret"] != "")
return creds, nil
}
// Use legacy fields directly
creds["type"] = tc.SourceType
creds["host"] = tc.SourceHost
creds["port"] = tc.SourcePort
creds["username"] = tc.SourceUser
creds["key_file"] = tc.SourceKeyFile
// Handle S3 fields
creds["bucket"] = tc.SourceBucket
creds["region"] = tc.SourceRegion
creds["access_key"] = tc.SourceAccessKey
creds["endpoint"] = tc.SourceEndpoint
// Handle SMB fields
creds["share"] = tc.SourceShare
creds["domain"] = tc.SourceDomain
// Handle FTP fields
if tc.SourcePassiveMode != nil {
creds["passive_mode"] = *tc.SourcePassiveMode
}
// Handle OAuth fields
creds["client_id"] = tc.SourceClientID
creds["drive_id"] = tc.SourceDriveID
creds["team_drive"] = tc.SourceTeamDrive
if tc.SourceReadOnly != nil {
creds["read_only"] = *tc.SourceReadOnly
}
creds["start_year"] = tc.SourceStartYear
if tc.SourceIncludeArchived != nil {
creds["include_archived"] = *tc.SourceIncludeArchived
}
if tc.UseBuiltinAuthSource != nil {
creds["use_builtin_auth"] = *tc.UseBuiltinAuthSource
}
// Handle temporary form fields and their encrypted counterparts
if tc.SourcePassword != "" {
creds["password"] = tc.SourcePassword
}
if tc.SourceSecretKey != "" {
creds["secret_key"] = tc.SourceSecretKey
}
if tc.SourceClientSecret != "" {
creds["client_secret"] = tc.SourceClientSecret
}
// If we have a db interface, try to encrypt any sensitive fields
if db != nil {
switch dbImpl := db.(type) {
case *DB:
// Handle encrypted fields if they exist in the database
if tc.SourcePassword != "" {
if encrypted, err := dbImpl.EncryptCredential(tc.SourcePassword); err == nil {
creds["encrypted_password"] = encrypted
}
}
if tc.SourceSecretKey != "" {
if encrypted, err := dbImpl.EncryptCredential(tc.SourceSecretKey); err == nil {
creds["encrypted_secret_key"] = encrypted
}
}
if tc.SourceClientSecret != "" {
if encrypted, err := dbImpl.EncryptCredential(tc.SourceClientSecret); err == nil {
creds["encrypted_client_secret"] = encrypted
}
}
}
}
return creds, nil
}
// GetDestinationCredentials returns credential information for the destination, either directly or from the provider
// If db is provided, it will try to load the provider from the database if needed
func (tc *TransferConfig) GetDestinationCredentials(db interface{}) (map[string]interface{}, error) {
creds := make(map[string]interface{})
fmt.Printf("DEBUG GetDestCreds Start: ProviderID=%v, HasProvider=%v\n",
tc.DestinationProviderID,
tc.DestinationProvider != nil)
// If using provider reference and provider is loaded
if tc.IsUsingDestinationProviderReference() {
// Try to load provider from database if we have a valid ID but no provider
if tc.DestinationProvider == nil && db != nil {
// Try different types of DB interfaces to load the provider
switch dbImpl := db.(type) {
case *DB:
provider, err := dbImpl.GetStorageProvider(*tc.DestinationProviderID)
if err != nil {
return nil, fmt.Errorf("failed to load destination provider (ID %d): %w", *tc.DestinationProviderID, err)
}
tc.DestinationProvider = provider
case interface {
GetStorageProvider(id uint) (*StorageProvider, error)
}:
provider, err := dbImpl.GetStorageProvider(*tc.DestinationProviderID)
if err != nil {
return nil, fmt.Errorf("failed to load destination provider (ID %d): %w", *tc.DestinationProviderID, err)
}
tc.DestinationProvider = provider
default:
return nil, fmt.Errorf("destination provider not loaded and db interface cannot load providers")
}
}
// If we still don't have a provider or it has no ID, return error
if tc.DestinationProvider == nil || tc.DestinationProvider.ID == 0 {
return nil, fmt.Errorf("failed to load valid destination provider (ID %d)", *tc.DestinationProviderID)
}
if tc.DestinationProvider != nil {
fmt.Printf("DEBUG Provider Details:\n"+
" ID: %v\n"+
" Type: %v\n"+
" Host: %v\n"+
" Port: %v\n"+
" Username: %v\n"+
" HasEncryptedPassword: %v\n"+
" HasKeyFile: %v\n"+
" HasSecretKey: %v\n"+
" HasClientSecret: %v\n"+
" HasRefreshToken: %v\n",
tc.DestinationProvider.ID,
tc.DestinationProvider.Type,
tc.DestinationProvider.Host,
tc.DestinationProvider.Port,
tc.DestinationProvider.Username,
tc.DestinationProvider.EncryptedPassword != "",
tc.DestinationProvider.KeyFile != "",
tc.DestinationProvider.EncryptedSecretKey != "",
tc.DestinationProvider.EncryptedClientSecret != "",
tc.DestinationProvider.EncryptedRefreshToken != "")
}
// Copy credentials from provider
creds["type"] = tc.DestinationProvider.Type
creds["host"] = tc.DestinationProvider.Host
creds["port"] = tc.DestinationProvider.Port
creds["username"] = tc.DestinationProvider.Username
creds["encrypted_password"] = tc.DestinationProvider.EncryptedPassword
creds["key_file"] = tc.DestinationProvider.KeyFile
// Handle S3 fields
creds["bucket"] = tc.DestinationProvider.Bucket
creds["region"] = tc.DestinationProvider.Region
creds["access_key"] = tc.DestinationProvider.AccessKey
creds["encrypted_secret_key"] = tc.DestinationProvider.EncryptedSecretKey
creds["endpoint"] = tc.DestinationProvider.Endpoint
// Handle SMB fields
creds["share"] = tc.DestinationProvider.Share
creds["domain"] = tc.DestinationProvider.Domain
// Handle FTP fields
if tc.DestinationProvider.PassiveMode != nil {
creds["passive_mode"] = *tc.DestinationProvider.PassiveMode
}
// Handle OAuth fields
creds["client_id"] = tc.DestinationProvider.ClientID
creds["encrypted_client_secret"] = tc.DestinationProvider.EncryptedClientSecret
creds["encrypted_refresh_token"] = tc.DestinationProvider.EncryptedRefreshToken
creds["drive_id"] = tc.DestinationProvider.DriveID
creds["team_drive"] = tc.DestinationProvider.TeamDrive
if tc.DestinationProvider.ReadOnly != nil {
creds["read_only"] = *tc.DestinationProvider.ReadOnly
}
creds["start_year"] = tc.DestinationProvider.StartYear
if tc.DestinationProvider.IncludeArchived != nil {
creds["include_archived"] = *tc.DestinationProvider.IncludeArchived
}
if tc.DestinationProvider.UseBuiltinAuth != nil {
creds["use_builtin_auth"] = *tc.DestinationProvider.UseBuiltinAuth
}
if tc.DestinationProvider.Authenticated != nil {
creds["authenticated"] = *tc.DestinationProvider.Authenticated
}
fmt.Printf("DEBUG Final Provider Creds:\n"+
" type: %v\n"+
" host: %v\n"+
" port: %v\n"+
" username: %v\n"+
" has_encrypted_password: %v\n"+
" has_key_file: %v\n"+
" has_encrypted_secret_key: %v\n"+
" has_encrypted_client_secret: %v\n",
creds["type"],
creds["host"],
creds["port"],
creds["username"],
creds["encrypted_password"] != "",
creds["key_file"] != "",
creds["encrypted_secret_key"] != "",
creds["encrypted_client_secret"] != "")
return creds, nil
}
// Use legacy fields directly
creds["type"] = tc.DestinationType
creds["host"] = tc.DestHost
creds["port"] = tc.DestPort
creds["username"] = tc.DestUser
creds["key_file"] = tc.DestKeyFile
// Handle S3 fields
creds["bucket"] = tc.DestBucket
creds["region"] = tc.DestRegion
creds["access_key"] = tc.DestAccessKey
creds["endpoint"] = tc.DestEndpoint
// Handle SMB fields
creds["share"] = tc.DestShare
creds["domain"] = tc.DestDomain
// Handle FTP fields
if tc.DestPassiveMode != nil {
creds["passive_mode"] = *tc.DestPassiveMode
}
// Handle OAuth fields
creds["client_id"] = tc.DestClientID
creds["drive_id"] = tc.DestDriveID
creds["team_drive"] = tc.DestTeamDrive
if tc.DestReadOnly != nil {
creds["read_only"] = *tc.DestReadOnly
}
creds["start_year"] = tc.DestStartYear
if tc.DestIncludeArchived != nil {
creds["include_archived"] = *tc.DestIncludeArchived
}
if tc.UseBuiltinAuthDest != nil {
creds["use_builtin_auth"] = *tc.UseBuiltinAuthDest
}
// Handle temporary form fields and their encrypted counterparts
if tc.DestPassword != "" {
creds["password"] = tc.DestPassword
}
if tc.DestSecretKey != "" {
creds["secret_key"] = tc.DestSecretKey
}
if tc.DestClientSecret != "" {
creds["client_secret"] = tc.DestClientSecret
}
// If we have a db interface, try to encrypt any sensitive fields
if db != nil {
switch dbImpl := db.(type) {
case *DB:
// Handle encrypted fields if they exist in the database
if tc.DestPassword != "" {
if encrypted, err := dbImpl.EncryptCredential(tc.DestPassword); err == nil {
creds["encrypted_password"] = encrypted
}
}
if tc.DestSecretKey != "" {
if encrypted, err := dbImpl.EncryptCredential(tc.DestSecretKey); err == nil {
creds["encrypted_secret_key"] = encrypted
}
}
if tc.DestClientSecret != "" {
if encrypted, err := dbImpl.EncryptCredential(tc.DestClientSecret); err == nil {
creds["encrypted_client_secret"] = encrypted
}
}
}
}
return creds, nil
}
@@ -0,0 +1,367 @@
package db_test
import (
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/starfleetcptn/gomft/internal/db"
"github.com/stretchr/testify/assert"
"gorm.io/gorm"
)
// setupTransferConfigTestDB sets up a SQLite in-memory database for testing
func setupTransferConfigTestDB(t *testing.T) *db.DB {
testDB, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
if err != nil {
t.Fatalf("Failed to open in-memory SQLite database: %v", err)
}
// Create tables
err = testDB.AutoMigrate(&db.StorageProvider{}, &db.TransferConfig{}, &db.User{})
if err != nil {
t.Fatalf("Failed to migrate tables: %v", err)
}
// Create test user
user := &db.User{
Email: "test@example.com",
PasswordHash: "hashedpassword",
}
err = testDB.Create(user).Error
if err != nil {
t.Fatalf("Failed to create test user: %v", err)
}
return &db.DB{DB: testDB}
}
// cleanupTransferConfigTestDB cleans up the test database
func cleanupTransferConfigTestDB(t *testing.T, testDB *gorm.DB) {
sqlDB, err := testDB.DB()
if err != nil {
t.Fatalf("Failed to get SQL DB: %v", err)
}
sqlDB.Close()
}
// TestTransferConfigWithProviderReferences tests the TransferConfig with StorageProvider references
func TestTransferConfigWithProviderReferences(t *testing.T) {
testDB := setupTransferConfigTestDB(t)
defer cleanupTransferConfigTestDB(t, testDB.DB)
// Create test storage providers
sourceProvider := &db.StorageProvider{
Name: "Test Source SFTP",
Type: db.ProviderTypeSFTP,
Host: "source.example.com",
Port: 22,
Username: "sourceuser",
EncryptedPassword: "encrypted_password_source",
CreatedBy: 1,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
destProvider := &db.StorageProvider{
Name: "Test Destination S3",
Type: db.ProviderTypeS3,
AccessKey: "destkey",
EncryptedSecretKey: "encrypted_secret_key_dest",
Region: "us-west-1",
CreatedBy: 1,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Save providers to database
err := testDB.CreateStorageProvider(sourceProvider)
assert.NoError(t, err, "Failed to create source provider")
err = testDB.CreateStorageProvider(destProvider)
assert.NoError(t, err, "Failed to create destination provider")
// Create a transfer config with provider references
config := &db.TransferConfig{
Name: "Test Config with Provider References",
SourcePath: "/source/path",
DestinationPath: "/dest/path",
CreatedBy: 1,
SourceType: string(db.ProviderTypeSFTP), // Set for compatibility
DestinationType: string(db.ProviderTypeS3), // Set for compatibility
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Set provider references
config.SetSourceProvider(sourceProvider)
config.SetDestinationProvider(destProvider)
// Save to database
err = testDB.Create(config).Error
assert.NoError(t, err, "Failed to create transfer config")
// Test IsUsingProviderReferences methods
assert.True(t, config.IsUsingSourceProviderReference(), "Should be using source provider reference")
assert.True(t, config.IsUsingDestinationProviderReference(), "Should be using destination provider reference")
assert.True(t, config.IsUsingProviderReferences(), "Should be using provider references")
// Clear providers to test loading from DB
config.SourceProvider = nil
config.DestinationProvider = nil
// Test GetSourceCredentials
sourceCreds, err := config.GetSourceCredentials(testDB)
assert.NoError(t, err, "Failed to get source credentials")
assert.Equal(t, "source.example.com", sourceCreds["host"], "Source host mismatch")
assert.Equal(t, 22, sourceCreds["port"], "Source port mismatch")
assert.Equal(t, "sourceuser", sourceCreds["username"], "Source username mismatch")
assert.Equal(t, "encrypted_password_source", sourceCreds["encrypted_password"], "Source encrypted password mismatch")
// Test GetDestinationCredentials
destCreds, err := config.GetDestinationCredentials(testDB)
assert.NoError(t, err, "Failed to get destination credentials")
assert.Equal(t, "destkey", destCreds["access_key"], "Destination access key mismatch")
assert.Equal(t, "encrypted_secret_key_dest", destCreds["encrypted_secret_key"], "Destination encrypted secret key mismatch")
assert.Equal(t, "us-west-1", destCreds["region"], "Destination region mismatch")
// Test that providers were loaded
assert.NotNil(t, config.SourceProvider, "Source provider should be loaded")
assert.NotNil(t, config.DestinationProvider, "Destination provider should be loaded")
}
// TestTransferConfigWithoutProviderReferences tests the TransferConfig without StorageProvider references
func TestTransferConfigWithoutProviderReferences(t *testing.T) {
testDB := setupTransferConfigTestDB(t)
defer cleanupTransferConfigTestDB(t, testDB.DB)
// Create a transfer config without provider references (legacy mode)
config := &db.TransferConfig{
Name: "Test Config without Provider References",
SourceType: string(db.ProviderTypeSFTP),
SourceHost: "direct.example.com",
SourcePort: 2222,
SourceUser: "directuser",
SourcePassword: "directpass", // This would be in form only
SourcePath: "/direct/source",
DestinationType: string(db.ProviderTypeS3),
DestAccessKey: "directaccesskey",
DestSecretKey: "directsecretkey", // This would be in form only
DestRegion: "eu-central-1",
DestinationPath: "/direct/dest",
CreatedBy: 1,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Save to database
err := testDB.Create(config).Error
assert.NoError(t, err, "Failed to create direct transfer config")
// Test IsUsingProviderReferences methods
assert.False(t, config.IsUsingSourceProviderReference(), "Should not be using source provider reference")
assert.False(t, config.IsUsingDestinationProviderReference(), "Should not be using destination provider reference")
assert.False(t, config.IsUsingProviderReferences(), "Should not be using provider references")
// Test GetSourceCredentials
sourceCreds, err := config.GetSourceCredentials(testDB)
assert.NoError(t, err, "Failed to get direct source credentials")
assert.Equal(t, "direct.example.com", sourceCreds["host"], "Direct source host mismatch")
assert.Equal(t, 2222, sourceCreds["port"], "Direct source port mismatch")
assert.Equal(t, "directuser", sourceCreds["username"], "Direct source username mismatch")
assert.Equal(t, "directpass", sourceCreds["password"], "Direct source password mismatch")
// Test GetDestinationCredentials
destCreds, err := config.GetDestinationCredentials(testDB)
assert.NoError(t, err, "Failed to get direct destination credentials")
assert.Equal(t, "directaccesskey", destCreds["access_key"], "Direct destination access key mismatch")
assert.Equal(t, "directsecretkey", destCreds["secret_key"], "Direct destination secret key mismatch")
assert.Equal(t, "eu-central-1", destCreds["region"], "Direct destination region mismatch")
}
// TestTransferConfigMixedProviderReferences tests TransferConfig with mixed provider references
func TestTransferConfigMixedProviderReferences(t *testing.T) {
testDB := setupTransferConfigTestDB(t)
defer cleanupTransferConfigTestDB(t, testDB.DB)
// Create test storage provider for source only
sourceProvider := &db.StorageProvider{
Name: "Test Mixed Source",
Type: db.ProviderTypeFTP,
Host: "mixed-source.example.com",
Port: 21,
Username: "mixeduser",
EncryptedPassword: "encrypted_password_mixed",
CreatedBy: 1,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Save provider to database
err := testDB.CreateStorageProvider(sourceProvider)
assert.NoError(t, err, "Failed to create mixed source provider")
// Create a transfer config with mixed provider references
config := &db.TransferConfig{
Name: "Test Config with Mixed Provider References",
SourcePath: "/mixed/source",
DestinationType: string(db.ProviderTypeS3),
DestAccessKey: "mixedaccesskey",
DestSecretKey: "mixedsecretkey", // This would be in form only
DestRegion: "ap-northeast-1",
DestinationPath: "/mixed/dest",
CreatedBy: 1,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Set source provider reference only
config.SetSourceProvider(sourceProvider)
// Save to database
err = testDB.Create(config).Error
assert.NoError(t, err, "Failed to create mixed transfer config")
// Test reference methods
assert.True(t, config.IsUsingSourceProviderReference(), "Should be using source provider reference")
assert.False(t, config.IsUsingDestinationProviderReference(), "Should not be using destination provider reference")
assert.False(t, config.IsUsingProviderReferences(), "Should not be using both provider references")
// Clear provider to test loading from DB
config.SourceProvider = nil
// Test GetSourceCredentials
sourceCreds, err := config.GetSourceCredentials(testDB)
assert.NoError(t, err, "Failed to get mixed source credentials")
assert.Equal(t, "mixed-source.example.com", sourceCreds["host"], "Mixed source host mismatch")
assert.Equal(t, 21, sourceCreds["port"], "Mixed source port mismatch")
assert.Equal(t, "mixeduser", sourceCreds["username"], "Mixed source username mismatch")
assert.Equal(t, "encrypted_password_mixed", sourceCreds["encrypted_password"], "Mixed source encrypted password mismatch")
// Test GetDestinationCredentials
destCreds, err := config.GetDestinationCredentials(testDB)
assert.NoError(t, err, "Failed to get mixed destination credentials")
assert.Equal(t, "mixedaccesskey", destCreds["access_key"], "Mixed destination access key mismatch")
assert.Equal(t, "mixedsecretkey", destCreds["secret_key"], "Mixed destination secret key mismatch")
assert.Equal(t, "ap-northeast-1", destCreds["region"], "Mixed destination region mismatch")
// Test that source provider was loaded
assert.NotNil(t, config.SourceProvider, "Source provider should be loaded")
}
// TestTransferConfigNonExistentProviderReferences tests error handling for non-existent provider references
func TestTransferConfigNonExistentProviderReferences(t *testing.T) {
testDB := setupTransferConfigTestDB(t)
defer cleanupTransferConfigTestDB(t, testDB.DB)
// Create uint pointers for provider IDs
sourceProviderID := uint(999)
destProviderID := uint(888)
// Create a transfer config with references to non-existent providers
config := &db.TransferConfig{
Name: "Test Config with Non-existent Provider References",
SourcePath: "/source/path",
DestinationPath: "/dest/path",
CreatedBy: 1,
SourceType: string(db.ProviderTypeSFTP),
DestinationType: string(db.ProviderTypeS3),
SourceProviderID: &sourceProviderID, // Use pointer to uint
DestinationProviderID: &destProviderID, // Use pointer to uint
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Save to database
err := testDB.Create(config).Error
assert.NoError(t, err, "Failed to create transfer config with non-existent provider references")
// Test GetSourceCredentials - should return error for non-existent provider
sourceCreds, err := config.GetSourceCredentials(testDB)
assert.Error(t, err, "Should get error for non-existent source provider")
assert.Nil(t, sourceCreds, "Source credentials should be nil for non-existent provider")
assert.Contains(t, err.Error(), "record not found", "Error should mention record not found")
// Test GetDestinationCredentials - should return error for non-existent provider
destCreds, err := config.GetDestinationCredentials(testDB)
assert.Error(t, err, "Should get error for non-existent destination provider")
assert.Nil(t, destCreds, "Destination credentials should be nil for non-existent provider")
assert.Contains(t, err.Error(), "record not found", "Error should mention record not found")
}
// TestTransferConfigIncompatibleProviderTypes tests behavior when provider types don't match config types
func TestTransferConfigIncompatibleProviderTypes(t *testing.T) {
testDB := setupTransferConfigTestDB(t)
defer cleanupTransferConfigTestDB(t, testDB.DB)
// Create test storage providers
sourceProvider := &db.StorageProvider{
Name: "S3 Source",
Type: db.ProviderTypeS3,
AccessKey: "sourcekey",
EncryptedSecretKey: "encrypted_secret_key_source",
Region: "us-east-1",
CreatedBy: 1,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
destProvider := &db.StorageProvider{
Name: "FTP Destination",
Type: db.ProviderTypeFTP,
Host: "dest.example.com",
Port: 21,
Username: "destuser",
EncryptedPassword: "encrypted_password_dest",
CreatedBy: 1,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Save providers to database
err := testDB.CreateStorageProvider(sourceProvider)
assert.NoError(t, err, "Failed to create source provider")
err = testDB.CreateStorageProvider(destProvider)
assert.NoError(t, err, "Failed to create destination provider")
// Create a transfer config with incompatible type declarations
config := &db.TransferConfig{
Name: "Test Config with Incompatible Types",
SourcePath: "/source/path",
DestinationPath: "/dest/path",
CreatedBy: 1,
SourceType: "sftp", // This is incompatible with the S3 provider
DestinationType: "s3", // This is incompatible with the FTP provider
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Set provider references
config.SetSourceProvider(sourceProvider)
config.SetDestinationProvider(destProvider)
// Save to database
err = testDB.Create(config).Error
assert.NoError(t, err, "Failed to create transfer config with incompatible types")
// Test GetCredentials methods
sourceCreds, err := config.GetSourceCredentials(testDB)
assert.NoError(t, err, "Should still get credentials despite type mismatch")
// Verify we can still get credentials from the provider despite type mismatch
assert.Equal(t, "sourcekey", sourceCreds["access_key"], "Should get correct credentials from provider despite type mismatch")
// The config's type is not automatically updated to match the provider
// Instead, it remains as what was explicitly set
assert.Equal(t, "sftp", config.SourceType, "Source type should remain as explicitly set")
destCreds, err := config.GetDestinationCredentials(testDB)
assert.NoError(t, err, "Should still get credentials despite type mismatch")
// Verify we can still get credentials from the provider despite type mismatch
assert.Equal(t, "destuser", destCreds["username"], "Should get correct credentials from provider despite type mismatch")
// The config's type is not automatically updated to match the provider
assert.Equal(t, "s3", config.DestinationType, "Destination type should remain as explicitly set")
}
+518 -53
View File
@@ -9,6 +9,8 @@ import (
"regexp"
"strconv"
"strings"
"github.com/starfleetcptn/gomft/internal/encryption"
)
// --- TransferConfig Store Methods ---
@@ -21,14 +23,14 @@ func (db *DB) CreateTransferConfig(config *TransferConfig) error {
// GetTransferConfigs retrieves all transfer configs for a user
func (db *DB) GetTransferConfigs(userID uint) ([]TransferConfig, error) {
var configs []TransferConfig
err := db.Where("created_by = ?", userID).Find(&configs).Error
err := db.Preload("SourceProvider").Preload("DestinationProvider").Where("created_by = ?", userID).Find(&configs).Error
return configs, err
}
// GetTransferConfig retrieves a single transfer config by ID
func (db *DB) GetTransferConfig(id uint) (*TransferConfig, error) {
var config TransferConfig
err := db.First(&config, id).Error
err := db.Preload("SourceProvider").Preload("DestinationProvider").First(&config, id).Error
if err != nil {
return nil, err
}
@@ -88,25 +90,103 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
rclonePath = "rclone"
}
// Ensure providers are loaded if using references
if config.IsUsingSourceProviderReference() && config.SourceProvider == nil {
provider, err := db.GetStorageProvider(*config.SourceProviderID)
if err != nil {
return fmt.Errorf("failed to load source provider (ID %d): %v", *config.SourceProviderID, err)
}
config.SetSourceProvider(provider)
}
if config.IsUsingDestinationProviderReference() && config.DestinationProvider == nil {
provider, err := db.GetStorageProvider(*config.DestinationProviderID)
if err != nil {
return fmt.Errorf("failed to load destination provider (ID %d): %v", *config.DestinationProviderID, err)
}
config.SetDestinationProvider(provider)
}
// If we have a destination provider but no ID or zero ID, fix it
if config.DestinationProvider != nil && (config.DestinationProviderID == nil || *config.DestinationProviderID == 0) {
config.SetDestinationProvider(config.DestinationProvider)
}
// If we have an ID but no provider, load it
if config.DestinationProviderID != nil && *config.DestinationProviderID > 0 && config.DestinationProvider == nil {
provider, err := db.GetStorageProvider(*config.DestinationProviderID)
if err != nil {
return fmt.Errorf("failed to load destination provider (ID %d): %v", *config.DestinationProviderID, err)
}
config.SetDestinationProvider(provider)
}
// Double check that everything is synchronized
if config.IsUsingDestinationProviderReference() {
if config.DestinationProvider == nil {
return fmt.Errorf("destination provider reference is set (ID %d) but provider is nil", *config.DestinationProviderID)
}
if config.DestinationProviderID == nil || *config.DestinationProviderID != config.DestinationProvider.ID {
config.SetDestinationProvider(config.DestinationProvider) // Re-sync the ID
}
}
// Get source credentials, either from provider or directly from config
sourceCredentials, err := config.GetSourceCredentials(db)
if err != nil {
return fmt.Errorf("failed to get source credentials: %v", err)
}
// Get source type either from provider or directly from config
sourceType := config.SourceType
if sourceTypeFromCreds, ok := sourceCredentials["type"].(StorageProviderType); ok {
sourceType = string(sourceTypeFromCreds)
} else if sourceTypeFromCreds, ok := sourceCredentials["type"].(string); ok {
sourceType = sourceTypeFromCreds
}
sourceName := fmt.Sprintf("source_%d", config.ID)
fmt.Printf("Generated source name: %s\n", sourceName)
fmt.Printf("Final source type being used: %s\n", sourceType)
// Generate rclone config using rclone CLI for source
switch config.SourceType {
switch sourceType {
case "sftp", "hetzner":
args := []string{
"config", "create", sourceName, "sftp",
"host", config.SourceHost,
"user", config.SourceUser,
"port", fmt.Sprintf("%d", config.SourcePort),
"host", getStringValue(sourceCredentials, "host", config.SourceHost),
"user", getStringValue(sourceCredentials, "username", config.SourceUser),
"port", fmt.Sprintf("%d", getIntValue(sourceCredentials, "port", config.SourcePort)),
"--non-interactive",
"--config", configPath,
"--log-level", "ERROR",
}
// First try to get password from direct form input (transient)
password := ""
if config.SourcePassword != "" {
args = append(args, "pass", config.SourcePassword)
password = config.SourcePassword
} else if encryptedPwd, ok := sourceCredentials["encrypted_password"].(string); ok && encryptedPwd != "" {
// For provider references, get the decrypted password
decryptedPwd, err := db.DecryptCredential(encryptedPwd)
if err != nil {
return fmt.Errorf("failed to decrypt source password: %v", err)
}
password = decryptedPwd
} else if pwVal, ok := sourceCredentials["password"].(string); ok && pwVal != "" {
// For backward compatibility
password = pwVal
}
if config.SourceKeyFile != "" {
args = append(args, "key_file", config.SourceKeyFile)
if password != "" {
args = append(args, "pass", password)
}
keyFile := getStringValue(sourceCredentials, "key_file", config.SourceKeyFile)
if keyFile != "" {
args = append(args, "key_file", keyFile)
}
cmd := exec.Command(rclonePath, args...)
if output, err := cmd.CombinedOutput(); err != nil {
return fmt.Errorf("failed to create source config (sftp): %v\nOutput: %s", err, output)
@@ -116,16 +196,49 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
"config", "create", sourceName, "s3",
"provider", "AWS", // Assuming AWS provider, adjust if needed
"env_auth", "false",
"access_key_id", config.SourceAccessKey,
"secret_access_key", config.SourceSecretKey,
"region", config.SourceRegion,
"--non-interactive",
"--config", configPath,
"--log-level", "ERROR",
}
if config.SourceEndpoint != "" {
args = append(args, "endpoint", config.SourceEndpoint)
// Handle access key
accessKey := getStringValue(sourceCredentials, "access_key", config.SourceAccessKey)
if accessKey != "" {
args = append(args, "access_key_id", accessKey)
}
// Handle secret key with proper decryption if from provider
secretKey := ""
if config.SourceSecretKey != "" {
// Direct input from form (transient)
secretKey = config.SourceSecretKey
} else if encryptedSecret, ok := sourceCredentials["encrypted_secret_key"].(string); ok && encryptedSecret != "" {
// Provider reference with encrypted secret
decryptedSecret, err := db.DecryptCredential(encryptedSecret)
if err != nil {
return fmt.Errorf("failed to decrypt source secret key: %v", err)
}
secretKey = decryptedSecret
} else if secretVal, ok := sourceCredentials["secret_key"].(string); ok && secretVal != "" {
// Backward compatibility
secretKey = secretVal
}
if secretKey != "" {
args = append(args, "secret_access_key", secretKey)
}
// Add region
region := getStringValue(sourceCredentials, "region", config.SourceRegion)
if region != "" {
args = append(args, "region", region)
}
endpoint := getStringValue(sourceCredentials, "endpoint", config.SourceEndpoint)
if endpoint != "" {
args = append(args, "endpoint", endpoint)
}
cmd := exec.Command(rclonePath, args...)
if output, err := cmd.CombinedOutput(); err != nil {
return fmt.Errorf("failed to create source config (s3): %v\nOutput: %s", err, output)
@@ -135,17 +248,19 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
"config", "create", sourceName, "s3",
"provider", "Wasabi",
"env_auth", "false",
"access_key_id", config.SourceAccessKey,
"secret_access_key", config.SourceSecretKey,
"region", config.SourceRegion,
"access_key_id", getStringValue(sourceCredentials, "access_key", config.SourceAccessKey),
"secret_access_key", getStringOrDefault(sourceCredentials, "secret_key", config.SourceSecretKey),
"region", getStringValue(sourceCredentials, "region", config.SourceRegion),
"--non-interactive",
"--config", configPath,
"--log-level", "ERROR",
}
endpoint := config.SourceEndpoint
endpoint := getStringValue(sourceCredentials, "endpoint", config.SourceEndpoint)
if endpoint == "" {
endpoint = "s3.wasabisys.com"
}
args = append(args, "endpoint", endpoint)
cmd := exec.Command(rclonePath, args...)
if output, err := cmd.CombinedOutput(); err != nil {
@@ -172,16 +287,16 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
"config", "create", sourceName, "s3",
"provider", "Minio",
"env_auth", "false",
"access_key_id", config.SourceAccessKey,
"secret_access_key", config.SourceSecretKey,
"endpoint", config.SourceEndpoint,
"access_key_id", getStringValue(sourceCredentials, "access_key", config.SourceAccessKey),
"secret_access_key", getStringOrDefault(sourceCredentials, "secret_key", config.SourceSecretKey),
"endpoint", getStringValue(sourceCredentials, "endpoint", config.SourceEndpoint),
"--non-interactive",
"--config", configPath,
"--log-level", "ERROR",
}
// Add region if specified
if config.SourceRegion != "" {
args = append(args, "region", config.SourceRegion)
if getStringValue(sourceCredentials, "region", config.SourceRegion) != "" {
args = append(args, "region", getStringValue(sourceCredentials, "region", config.SourceRegion))
}
cmd := exec.Command(rclonePath, args...)
if output, err := cmd.CombinedOutput(); err != nil {
@@ -228,7 +343,7 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
if len(output) > 0 {
errorMsg += fmt.Sprintf("\nOutput: %s", output)
}
return fmt.Errorf(errorMsg)
return fmt.Errorf("%v", errorMsg)
}
case "local":
// For local source, ensure the section exists but might not need specific rclone config create
@@ -242,25 +357,58 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
}
// Get destination credentials, either from provider or directly from config
destCredentials, err := config.GetDestinationCredentials(db)
if err != nil {
return fmt.Errorf("failed to get destination credentials: %v", err)
}
// Get destination type either from provider or directly from config
destType := config.DestinationType
if destTypeFromCreds, ok := destCredentials["type"].(StorageProviderType); ok {
destType = string(destTypeFromCreds)
} else if destTypeFromCreds, ok := destCredentials["type"].(string); ok {
destType = destTypeFromCreds
}
destName := fmt.Sprintf("dest_%d", config.ID)
// Generate rclone config using rclone CLI for destination
switch config.DestinationType {
switch destType {
case "sftp", "hetzner":
args := []string{
"config", "create", destName, "sftp",
"host", config.DestHost,
"user", config.DestUser,
"port", fmt.Sprintf("%d", config.DestPort),
"host", getStringValue(destCredentials, "host", config.DestHost),
"user", getStringValue(destCredentials, "username", config.DestUser),
"port", fmt.Sprintf("%d", getIntValue(destCredentials, "port", config.DestPort)),
"--non-interactive",
"--config", configPath,
"--log-level", "ERROR",
}
password := ""
if config.DestPassword != "" {
args = append(args, "pass", config.DestPassword)
password = config.DestPassword
} else if encryptedPwd, ok := destCredentials["encrypted_password"].(string); ok && encryptedPwd != "" {
// For provider references, get the decrypted password
decryptedPwd, err := db.DecryptCredential(encryptedPwd)
if err != nil {
return fmt.Errorf("failed to decrypt destination password: %v", err)
}
password = decryptedPwd
} else if pwVal, ok := destCredentials["password"].(string); ok && pwVal != "" {
// For backward compatibility
password = pwVal
}
if config.DestKeyFile != "" {
args = append(args, "key_file", config.DestKeyFile)
if password != "" {
args = append(args, "pass", password)
}
keyFile := getStringValue(destCredentials, "key_file", config.DestKeyFile)
if keyFile != "" {
args = append(args, "key_file", keyFile)
}
cmd := exec.Command(rclonePath, args...)
if output, err := cmd.CombinedOutput(); err != nil {
return fmt.Errorf("failed to create destination config (sftp): %v\nOutput: %s", err, output)
@@ -270,16 +418,19 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
"config", "create", destName, "s3",
"provider", "AWS", // Assuming AWS provider
"env_auth", "false",
"access_key_id", config.DestAccessKey,
"secret_access_key", config.DestSecretKey,
"region", config.DestRegion,
"access_key_id", getStringValue(destCredentials, "access_key", config.DestAccessKey),
"secret_access_key", getStringOrDefault(destCredentials, "secret_key", config.DestSecretKey),
"region", getStringValue(destCredentials, "region", config.DestRegion),
"--non-interactive",
"--config", configPath,
"--log-level", "ERROR",
}
if config.DestEndpoint != "" {
args = append(args, "endpoint", config.DestEndpoint)
endpoint := getStringValue(destCredentials, "endpoint", config.DestEndpoint)
if endpoint != "" {
args = append(args, "endpoint", endpoint)
}
cmd := exec.Command(rclonePath, args...)
if output, err := cmd.CombinedOutput(); err != nil {
return fmt.Errorf("failed to create destination config (s3): %v\nOutput: %s", err, output)
@@ -289,17 +440,19 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
"config", "create", destName, "s3",
"provider", "Wasabi",
"env_auth", "false",
"access_key_id", config.DestAccessKey,
"secret_access_key", config.DestSecretKey,
"region", config.DestRegion,
"access_key_id", getStringValue(destCredentials, "access_key", config.DestAccessKey),
"secret_access_key", getStringOrDefault(destCredentials, "secret_key", config.DestSecretKey),
"region", getStringValue(destCredentials, "region", config.DestRegion),
"--non-interactive",
"--config", configPath,
"--log-level", "ERROR",
}
endpoint := config.DestEndpoint
endpoint := getStringValue(destCredentials, "endpoint", config.DestEndpoint)
if endpoint == "" {
endpoint = "s3.wasabisys.com"
}
args = append(args, "endpoint", endpoint)
cmd := exec.Command(rclonePath, args...)
if output, err := cmd.CombinedOutput(); err != nil {
@@ -308,15 +461,43 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
case "b2":
args := []string{
"config", "create", destName, "b2",
"account", config.DestAccessKey, // B2 Account ID
"key", config.DestSecretKey, // B2 Application Key
"--non-interactive",
"--config", configPath,
"--log-level", "ERROR",
}
if config.DestEndpoint != "" {
args = append(args, "endpoint", config.DestEndpoint)
// Handle account ID (access key)
accountID := getStringValue(destCredentials, "access_key", config.DestAccessKey)
if accountID != "" {
args = append(args, "account", accountID)
}
// Handle application key (secret key) with proper decryption if from provider
appKey := ""
if config.DestSecretKey != "" {
// Direct input from form (transient)
appKey = config.DestSecretKey
} else if encryptedSecret, ok := destCredentials["encrypted_secret_key"].(string); ok && encryptedSecret != "" {
// Provider reference with encrypted secret
decryptedSecret, err := db.DecryptCredential(encryptedSecret)
if err != nil {
return fmt.Errorf("failed to decrypt destination secret key: %v", err)
}
appKey = decryptedSecret
} else if secretVal, ok := destCredentials["secret_key"].(string); ok && secretVal != "" {
// Backward compatibility
appKey = secretVal
}
if appKey != "" {
args = append(args, "key", appKey)
}
endpoint := getStringValue(destCredentials, "endpoint", config.DestEndpoint)
if endpoint != "" {
args = append(args, "endpoint", endpoint)
}
cmd := exec.Command(rclonePath, args...)
if output, err := cmd.CombinedOutput(); err != nil {
return fmt.Errorf("failed to create destination config (b2): %v\nOutput: %s", err, output)
@@ -326,16 +507,16 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
"config", "create", destName, "s3",
"provider", "Minio",
"env_auth", "false",
"access_key_id", config.DestAccessKey,
"secret_access_key", config.DestSecretKey,
"endpoint", config.DestEndpoint,
"access_key_id", getStringValue(destCredentials, "access_key", config.DestAccessKey),
"secret_access_key", getStringOrDefault(destCredentials, "secret_key", config.DestSecretKey),
"endpoint", getStringValue(destCredentials, "endpoint", config.DestEndpoint),
"--non-interactive",
"--config", configPath,
"--log-level", "ERROR",
}
// Add region if specified
if config.DestRegion != "" {
args = append(args, "region", config.DestRegion)
if getStringValue(destCredentials, "region", config.DestRegion) != "" {
args = append(args, "region", getStringValue(destCredentials, "region", config.DestRegion))
}
cmd := exec.Command(rclonePath, args...)
if output, err := cmd.CombinedOutput(); err != nil {
@@ -361,26 +542,47 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
if config.DestinationType == "nextcloud" {
vendor = "nextcloud"
webdavURL = fmt.Sprintf("%s/remote.php/dav/files/%s/", webdavURL, config.DestUser) // Corrected variable
webdavURL = fmt.Sprintf("%s/remote.php/dav/files/%s/", webdavURL, getStringValue(destCredentials, "username", config.DestUser)) // Corrected variable
}
args := []string{
"config", "create", destName, "webdav",
"url", webdavURL, // Use the parsed and reconstructed URL
"vendor", vendor,
"user", config.DestUser,
"pass", config.DestPassword, // rclone obscures this
"user", getStringValue(destCredentials, "username", config.DestUser),
"--non-interactive",
"--config", configPath,
"--log-level", "ERROR",
}
// Handle password with proper decryption if from provider
password := ""
if config.DestPassword != "" {
// Direct input from form (transient)
password = config.DestPassword
} else if encryptedPwd, ok := destCredentials["encrypted_password"].(string); ok && encryptedPwd != "" {
// Provider reference with encrypted password
decryptedPwd, err := db.DecryptCredential(encryptedPwd)
if err != nil {
return fmt.Errorf("failed to decrypt destination password: %v", err)
}
password = decryptedPwd
} else if pwVal, ok := destCredentials["password"].(string); ok && pwVal != "" {
// Backward compatibility
password = pwVal
}
if password != "" {
args = append(args, "pass", password)
}
cmd := exec.Command(rclonePath, args...)
if output, err := cmd.CombinedOutput(); err != nil {
errorMsg := fmt.Sprintf("failed to create destination config (%s): %v", config.DestinationType, err)
if len(output) > 0 {
errorMsg += fmt.Sprintf("\nOutput: %s", output)
}
return fmt.Errorf(errorMsg)
return fmt.Errorf("%v", errorMsg)
}
case "local":
// Append local config section
@@ -401,6 +603,31 @@ func (db *DB) GenerateRcloneConfig(config *TransferConfig) error {
return nil
}
// Helper functions to get values from credentials map
func getStringValue(creds map[string]interface{}, key, defaultValue string) string {
if val, ok := creds[key].(string); ok && val != "" {
return val
}
return defaultValue
}
func getStringOrDefault(creds map[string]interface{}, key, defaultValue string) string {
if defaultValue != "" {
return defaultValue // Prefer the value passed directly for sensitive fields
}
if val, ok := creds[key].(string); ok {
return val
}
return ""
}
func getIntValue(creds map[string]interface{}, key string, defaultValue int) int {
if val, ok := creds[key].(int); ok {
return val
}
return defaultValue
}
// StoreGoogleDriveToken stores the Google Drive auth token for a config
func (db *DB) StoreGoogleDriveToken(configIDStr string, token string) error {
configID, err := strconv.ParseUint(configIDStr, 10, 64)
@@ -420,6 +647,36 @@ func (db *DB) StoreGoogleDriveToken(configIDStr string, token string) error {
return fmt.Errorf("failed to update config: %v", err)
}
// Check if we're using a provider reference and update the provider instead
if config.IsUsingDestinationProviderReference() && config.DestinationProvider != nil &&
(config.DestinationProvider.Type == "gdrive" || config.DestinationProvider.Type == "gphotos") {
// Update the provider with the token
provider := config.DestinationProvider
provider.RefreshToken = token // Set the clear token temporarily
provider.SetAuthenticated(true)
// Update the provider in the database
if err := db.UpdateStorageProvider(provider); err != nil {
return fmt.Errorf("failed to update provider with token: %v", err)
}
// Continue with creating the rclone config file since this is still needed for transfers
} else if config.IsUsingSourceProviderReference() && config.SourceProvider != nil &&
(config.SourceProvider.Type == "gdrive" || config.SourceProvider.Type == "gphotos") {
// Update the provider with the token
provider := config.SourceProvider
provider.RefreshToken = token // Set the clear token temporarily
provider.SetAuthenticated(true)
// Update the provider in the database
if err := db.UpdateStorageProvider(provider); err != nil {
return fmt.Errorf("failed to update provider with token: %v", err)
}
// Continue with creating the rclone config file since this is still needed for transfers
}
// Legacy fallback for direct token storage in config file
configPath := db.GetConfigRclonePath(config)
existingConfig := ""
if _, err := os.Stat(configPath); err == nil {
@@ -660,3 +917,211 @@ func (db *DB) GetGDriveCredentialsFromConfig(config *TransferConfig) (string, st
}
return "", ""
}
// ConvertToProviderReferences converts a TransferConfig that uses embedded credentials
// to one that uses StorageProvider references.
func (db *DB) ConvertToProviderReferences(config *TransferConfig) error {
// Skip if already using both provider references
if config.IsUsingProviderReferences() {
return nil
}
tx := db.Begin()
if tx.Error != nil {
return fmt.Errorf("failed to start transaction: %v", tx.Error)
}
defer func() {
if r := recover(); r != nil {
tx.Rollback()
}
}()
// Convert source if needed
if !config.IsUsingSourceProviderReference() && config.SourceType != "" {
// Create new provider from source fields
provider := &StorageProvider{
Name: fmt.Sprintf("%s Source - %s", config.Name, config.SourceType),
Type: StorageProviderType(config.SourceType),
CreatedBy: config.CreatedBy,
// Copy all relevant source fields to provider fields
Host: config.SourceHost,
Port: config.SourcePort,
Username: config.SourceUser,
KeyFile: config.SourceKeyFile,
Bucket: config.SourceBucket,
Region: config.SourceRegion,
AccessKey: config.SourceAccessKey,
Share: config.SourceShare,
Domain: config.SourceDomain,
PassiveMode: config.SourcePassiveMode,
ClientID: config.SourceClientID,
DriveID: config.SourceDriveID,
TeamDrive: config.SourceTeamDrive,
ReadOnly: config.SourceReadOnly,
StartYear: config.SourceStartYear,
IncludeArchived: config.SourceIncludeArchived,
UseBuiltinAuth: config.UseBuiltinAuthSource,
}
// Handle fields that need encryption
if config.SourcePassword != "" {
encryptedPwd, err := db.EncryptCredential(config.SourcePassword)
if err != nil {
tx.Rollback()
return fmt.Errorf("failed to encrypt source password: %v", err)
}
provider.EncryptedPassword = encryptedPwd
}
if config.SourceSecretKey != "" {
encryptedSecret, err := db.EncryptCredential(config.SourceSecretKey)
if err != nil {
tx.Rollback()
return fmt.Errorf("failed to encrypt source secret key: %v", err)
}
provider.EncryptedSecretKey = encryptedSecret
}
if config.SourceClientSecret != "" {
encryptedClientSecret, err := db.EncryptCredential(config.SourceClientSecret)
if err != nil {
tx.Rollback()
return fmt.Errorf("failed to encrypt source client secret: %v", err)
}
provider.EncryptedClientSecret = encryptedClientSecret
}
// Save the new provider
if err := tx.Create(provider).Error; err != nil {
tx.Rollback()
return fmt.Errorf("failed to create source provider: %v", err)
}
// Update the config to reference the new provider
config.SetSourceProvider(provider)
}
// Convert destination if needed
if !config.IsUsingDestinationProviderReference() && config.DestinationType != "" {
// Create new provider from destination fields
provider := &StorageProvider{
Name: fmt.Sprintf("%s Destination - %s", config.Name, config.DestinationType),
Type: StorageProviderType(config.DestinationType),
CreatedBy: config.CreatedBy,
// Copy all relevant destination fields to provider fields
Host: config.DestHost,
Port: config.DestPort,
Username: config.DestUser,
KeyFile: config.DestKeyFile,
Bucket: config.DestBucket,
Region: config.DestRegion,
AccessKey: config.DestAccessKey,
Share: config.DestShare,
Domain: config.DestDomain,
PassiveMode: config.DestPassiveMode,
ClientID: config.DestClientID,
DriveID: config.DestDriveID,
TeamDrive: config.DestTeamDrive,
ReadOnly: config.DestReadOnly,
StartYear: config.DestStartYear,
IncludeArchived: config.DestIncludeArchived,
UseBuiltinAuth: config.UseBuiltinAuthDest,
}
// For Google Drive/Photos, carry over authentication status
if config.DestinationType == "gdrive" || config.DestinationType == "gphotos" {
provider.SetAuthenticated(config.GetGoogleAuthenticated())
}
// Handle fields that need encryption
if config.DestPassword != "" {
encryptedPwd, err := db.EncryptCredential(config.DestPassword)
if err != nil {
tx.Rollback()
return fmt.Errorf("failed to encrypt destination password: %v", err)
}
provider.EncryptedPassword = encryptedPwd
}
if config.DestSecretKey != "" {
encryptedSecret, err := db.EncryptCredential(config.DestSecretKey)
if err != nil {
tx.Rollback()
return fmt.Errorf("failed to encrypt destination secret key: %v", err)
}
provider.EncryptedSecretKey = encryptedSecret
}
if config.DestClientSecret != "" {
encryptedClientSecret, err := db.EncryptCredential(config.DestClientSecret)
if err != nil {
tx.Rollback()
return fmt.Errorf("failed to encrypt destination client secret: %v", err)
}
provider.EncryptedClientSecret = encryptedClientSecret
}
// Save the new provider
if err := tx.Create(provider).Error; err != nil {
tx.Rollback()
return fmt.Errorf("failed to create destination provider: %v", err)
}
// Update the config to reference the new provider
config.SetDestinationProvider(provider)
}
// Save the updated config
if err := tx.Save(config).Error; err != nil {
tx.Rollback()
return fmt.Errorf("failed to update config with provider references: %v", err)
}
return tx.Commit().Error
}
// EncryptCredential encrypts a sensitive credential value
func (db *DB) EncryptCredential(value string) (string, error) {
// Create a credential encryptor
encryptor, err := encryption.GetGlobalCredentialEncryptor()
if err != nil {
return "", fmt.Errorf("failed to get credential encryptor: %w", err)
}
// Encrypt the value using the generic credential type
encrypted, err := encryptor.Encrypt(value, encryption.TypeGeneric)
if err != nil {
return "", fmt.Errorf("failed to encrypt credential: %w", err)
}
return encrypted, nil
}
// DecryptCredential decrypts a sensitive credential value
func (db *DB) DecryptCredential(encryptedValue string) (string, error) {
// Create a credential encryptor
encryptor, err := encryption.GetGlobalCredentialEncryptor()
if err != nil {
return "", fmt.Errorf("failed to get credential encryptor: %w", err)
}
// Check if value is already encrypted with our prefix
if !encryptor.IsEncrypted(encryptedValue) {
// Handle legacy format (temporary backward compatibility)
if strings.HasPrefix(encryptedValue, "encrypted_") {
return strings.TrimPrefix(encryptedValue, "encrypted_"), nil
}
// Not encrypted, return as-is
return encryptedValue, nil
}
// Decrypt the value
decrypted, err := encryptor.Decrypt(encryptedValue)
if err != nil {
return "", fmt.Errorf("failed to decrypt credential: %w", err)
}
return decrypted, nil
}
// UpdateStorageProvider updates an existing storage provider