mirror of
https://github.com/StarFleetCPTN/GoMFT.git
synced 2026-10-11 00:37:28 +02:00
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:
1 parent
88d0ac815a
commit
31871bd16e
59 files changed
+15842
-594
No files matched your search
+62
-2
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
Reference in new issue
Block a user