feat: Implement storage provider management functionality

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

No files matched your search

+400
View File
@@ -0,0 +1,400 @@
package audit
import (
"encoding/json"
"fmt"
"io"
"os"
"sync"
"time"
"github.com/starfleetcptn/gomft/internal/encryption"
)
// EventType represents the type of encryption-related event
type EventType string
// Event types for encryption operations
const (
EventEncrypt EventType = "encrypt"
EventDecrypt EventType = "decrypt"
EventKeyAccess EventType = "key_access"
EventKeyRotation EventType = "key_rotation"
EventKeyGeneration EventType = "key_generation"
EventDecryptionFailure EventType = "decryption_failure"
EventEncryptionFailure EventType = "encryption_failure"
)
// SecurityLevel represents the severity/importance of an audit event
type SecurityLevel string
// Security levels for events
const (
LevelInfo SecurityLevel = "info"
LevelWarning SecurityLevel = "warning"
LevelAlert SecurityLevel = "alert"
LevelError SecurityLevel = "error"
)
// AuditEvent represents a single encryption-related security event
type AuditEvent struct {
Timestamp time.Time `json:"timestamp"`
EventType EventType `json:"event_type"`
Level SecurityLevel `json:"level"`
Operation string `json:"operation"`
FieldType string `json:"field_type,omitempty"`
ModelType string `json:"model_type,omitempty"`
Description string `json:"description"`
Success bool `json:"success"`
Error string `json:"error,omitempty"`
KeyVersion string `json:"key_version,omitempty"`
UserID uint `json:"user_id,omitempty"`
RemoteIP string `json:"remote_ip,omitempty"`
Duration int64 `json:"duration_ns,omitempty"` // Operation duration in nanoseconds
}
// SecurityAuditor is responsible for logging security-related events
type SecurityAuditor struct {
enabled bool
logWriter io.Writer
errorWriter io.Writer
mutex sync.Mutex
detailedMode bool
logFilePath string
errorFilePath string
}
// New creates a new SecurityAuditor with default configuration
func New() (*SecurityAuditor, error) {
return &SecurityAuditor{
enabled: true,
logWriter: os.Stdout, // Default to stdout for regular logs
errorWriter: os.Stderr, // Default to stderr for error logs
detailedMode: false,
}, nil
}
// NewWithFileLogging creates a new SecurityAuditor with file-based logging
func NewWithFileLogging(logFilePath, errorFilePath string) (*SecurityAuditor, error) {
logFile, err := os.OpenFile(logFilePath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
if err != nil {
return nil, fmt.Errorf("failed to open log file: %w", err)
}
var errorWriter io.Writer
if errorFilePath == logFilePath {
errorWriter = logFile
} else {
errorFile, err := os.OpenFile(errorFilePath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
if err != nil {
logFile.Close()
return nil, fmt.Errorf("failed to open error log file: %w", err)
}
errorWriter = errorFile
}
return &SecurityAuditor{
enabled: true,
logWriter: logFile,
errorWriter: errorWriter,
logFilePath: logFilePath,
errorFilePath: errorFilePath,
detailedMode: false,
}, nil
}
// Close properly closes any open resources
func (a *SecurityAuditor) Close() error {
a.mutex.Lock()
defer a.mutex.Unlock()
// Check if we need to close file writers
if closer, ok := a.logWriter.(io.Closer); ok {
if err := closer.Close(); err != nil {
return err
}
}
// Don't close errorWriter if it's the same as logWriter
if a.errorFilePath != a.logFilePath {
if closer, ok := a.errorWriter.(io.Closer); ok {
if err := closer.Close(); err != nil {
return err
}
}
}
return nil
}
// Enable turns on the auditor
func (a *SecurityAuditor) Enable() {
a.mutex.Lock()
defer a.mutex.Unlock()
a.enabled = true
}
// Disable turns off the auditor
func (a *SecurityAuditor) Disable() {
a.mutex.Lock()
defer a.mutex.Unlock()
a.enabled = false
}
// SetDetailedMode toggles detailed logging mode
func (a *SecurityAuditor) SetDetailedMode(detailed bool) {
a.mutex.Lock()
defer a.mutex.Unlock()
a.detailedMode = detailed
}
// IsEnabled returns whether auditing is enabled
func (a *SecurityAuditor) IsEnabled() bool {
a.mutex.Lock()
defer a.mutex.Unlock()
return a.enabled
}
// LogEvent records a security event to the audit log
func (a *SecurityAuditor) LogEvent(event AuditEvent) {
if !a.IsEnabled() {
return
}
a.mutex.Lock()
defer a.mutex.Unlock()
// Ensure timestamp is set
if event.Timestamp.IsZero() {
event.Timestamp = time.Now()
}
// Convert the event to JSON
jsonData, err := json.Marshal(event)
if err != nil {
fmt.Fprintf(a.errorWriter, "Error marshaling audit event: %v\n", err)
return
}
// Choose the right writer based on event level
writer := a.logWriter
if event.Level == LevelError || event.Level == LevelAlert {
writer = a.errorWriter
}
// Write to the appropriate log
fmt.Fprintln(writer, string(jsonData))
}
// LogEncryptionEvent logs an encryption operation event
func (a *SecurityAuditor) LogEncryptionEvent(operation string, fieldType, modelType string, success bool, err error, keyVersion string, userID uint, duration time.Duration) {
if !a.IsEnabled() {
return
}
event := AuditEvent{
Timestamp: time.Now(),
EventType: EventEncrypt,
Level: LevelInfo,
Operation: operation,
FieldType: fieldType,
ModelType: modelType,
Success: success,
KeyVersion: keyVersion,
UserID: userID,
Duration: duration.Nanoseconds(),
}
if !success {
event.EventType = EventEncryptionFailure
event.Level = LevelWarning
if err != nil {
event.Error = encryption.SanitizeError(err.Error())
}
}
a.LogEvent(event)
}
// LogDecryptionEvent logs a decryption operation event
func (a *SecurityAuditor) LogDecryptionEvent(operation string, fieldType, modelType string, success bool, err error, keyVersion string, userID uint, duration time.Duration) {
if !a.IsEnabled() {
return
}
event := AuditEvent{
Timestamp: time.Now(),
EventType: EventDecrypt,
Level: LevelInfo,
Operation: operation,
FieldType: fieldType,
ModelType: modelType,
Success: success,
KeyVersion: keyVersion,
UserID: userID,
Duration: duration.Nanoseconds(),
}
if !success {
event.EventType = EventDecryptionFailure
event.Level = LevelWarning
if err != nil {
event.Error = encryption.SanitizeError(err.Error())
}
}
a.LogEvent(event)
}
// LogKeyAccessEvent logs when an encryption key is accessed
func (a *SecurityAuditor) LogKeyAccessEvent(keyVersion string, success bool, err error, userID uint) {
if !a.IsEnabled() {
return
}
event := AuditEvent{
Timestamp: time.Now(),
EventType: EventKeyAccess,
Level: LevelInfo,
Operation: "key_access",
Success: success,
KeyVersion: keyVersion,
UserID: userID,
}
if !success {
event.Level = LevelAlert
if err != nil {
event.Error = encryption.SanitizeError(err.Error())
}
}
// Key access failures are security-critical and should be logged at a higher level
if !success {
event.Description = "Failed key access attempt"
}
a.LogEvent(event)
}
// LogKeyRotationEvent logs when encryption keys are rotated
func (a *SecurityAuditor) LogKeyRotationEvent(oldVersion, newVersion string, success bool, err error, userID uint) {
if !a.IsEnabled() {
return
}
event := AuditEvent{
Timestamp: time.Now(),
EventType: EventKeyRotation,
Level: LevelInfo,
Operation: "key_rotation",
Description: fmt.Sprintf("Key rotation from version %s to %s", oldVersion, newVersion),
Success: success,
KeyVersion: newVersion,
UserID: userID,
}
if !success {
event.Level = LevelError
if err != nil {
event.Error = encryption.SanitizeError(err.Error())
}
}
a.LogEvent(event)
}
// LogKeyRotationEventWithDescription logs when encryption keys are rotated with a custom description
func (a *SecurityAuditor) LogKeyRotationEventWithDescription(oldVersion, newVersion string, success bool, description string, userID uint) {
if !a.IsEnabled() {
return
}
event := AuditEvent{
Timestamp: time.Now(),
EventType: EventKeyRotation,
Level: LevelInfo,
Operation: "key_rotation",
Description: description,
Success: success,
KeyVersion: newVersion,
UserID: userID,
}
if !success {
event.Level = LevelError
}
a.LogEvent(event)
}
// LogKeyGenerationEvent logs when a new encryption key is generated
func (a *SecurityAuditor) LogKeyGenerationEvent(keyVersion string, success bool, err error, userID uint) {
if !a.IsEnabled() {
return
}
event := AuditEvent{
Timestamp: time.Now(),
EventType: EventKeyGeneration,
Level: LevelInfo,
Operation: "key_generation",
Description: "New encryption key generated",
Success: success,
KeyVersion: keyVersion,
UserID: userID,
}
if !success {
event.Level = LevelError
if err != nil {
event.Error = encryption.SanitizeError(err.Error())
}
}
a.LogEvent(event)
}
// global is the default security auditor instance
var global *SecurityAuditor
var globalOnce sync.Once
// GetGlobalAuditor returns the global security auditor instance
func GetGlobalAuditor() *SecurityAuditor {
globalOnce.Do(func() {
var err error
global, err = New()
if err != nil {
// Fall back to a disabled auditor if there's an error
global = &SecurityAuditor{enabled: false}
}
})
return global
}
// InitializeWithFileLogging initializes the global auditor with file logging
func InitializeWithFileLogging(logFilePath, errorFilePath string) error {
auditor, err := NewWithFileLogging(logFilePath, errorFilePath)
if err != nil {
return err
}
globalOnce.Do(func() {
global = auditor
})
// If global auditor was already initialized, replace it
if global != auditor {
if closer, ok := global.logWriter.(io.Closer); ok {
closer.Close()
}
if global.errorFilePath != global.logFilePath {
if closer, ok := global.errorWriter.(io.Closer); ok {
closer.Close()
}
}
global = auditor
}
return nil
}
@@ -0,0 +1,512 @@
package audit
import (
"context"
"fmt"
"reflect"
"runtime"
"strings"
"sync"
"time"
"github.com/starfleetcptn/gomft/internal/encryption"
"github.com/starfleetcptn/gomft/internal/encryption/keyrotation"
"gorm.io/gorm"
)
// RotationOptions contains configuration for the key rotation process
type RotationOptions struct {
// DryRun performs all operations but doesn't save changes to database
DryRun bool
// BatchSize sets the number of records to process in each batch
BatchSize int
// MaxErrors sets the threshold of errors before aborting
MaxErrors int
// Parallelism controls how many models are processed in parallel
Parallelism int
// Timeout specifies a maximum duration for the entire operation
Timeout time.Duration
// WorkerTimeout specifies maximum duration for a single batch
WorkerTimeout time.Duration
// ProgressCallback receives updates on rotation progress
ProgressCallback func(modelName string, processed, total int)
}
// RotationUtility provides comprehensive capabilities for rotating encryption keys
// across multiple database models with detailed auditing and progress tracking
type RotationUtility struct {
db *gorm.DB
oldService *encryption.EncryptionService
newService *encryption.EncryptionService
auditor *SecurityAuditor
monitor *SecurityMonitor
options RotationOptions
testingHooks map[string]func(interface{}) error
mu sync.Mutex
}
// NewRotationUtility creates a new RotationUtility
func NewRotationUtility(
db *gorm.DB,
oldService, newService *encryption.EncryptionService,
auditor *SecurityAuditor,
monitor *SecurityMonitor,
options RotationOptions,
) (*RotationUtility, error) {
if db == nil {
return nil, fmt.Errorf("database connection is required")
}
if oldService == nil {
return nil, fmt.Errorf("old encryption service is required")
}
if newService == nil {
return nil, fmt.Errorf("new encryption service is required")
}
if auditor == nil {
auditor = GetGlobalAuditor()
}
if monitor == nil {
monitor = NewSecurityMonitor(auditor)
}
// Set default options
if options.BatchSize <= 0 {
options.BatchSize = 100
}
if options.MaxErrors <= 0 {
options.MaxErrors = 50
}
if options.Parallelism <= 0 {
options.Parallelism = 1
}
if options.Timeout <= 0 {
options.Timeout = 24 * time.Hour // Default long timeout
}
if options.WorkerTimeout <= 0 {
options.WorkerTimeout = 30 * time.Minute
}
return &RotationUtility{
db: db,
oldService: oldService,
newService: newService,
auditor: auditor,
monitor: monitor,
options: options,
testingHooks: make(map[string]func(interface{}) error),
}, nil
}
// RegisterTestingHook registers a hook for testing purposes
func (r *RotationUtility) RegisterTestingHook(name string, hook func(interface{}) error) {
r.mu.Lock()
defer r.mu.Unlock()
r.testingHooks[name] = hook
}
// runHook runs a testing hook if it exists
func (r *RotationUtility) runHook(name string, data interface{}) error {
r.mu.Lock()
hook, exists := r.testingHooks[name]
r.mu.Unlock()
if exists && hook != nil {
return hook(data)
}
return nil
}
// RotateKeysForModels performs key rotation for multiple model types with detailed monitoring
func (r *RotationUtility) RotateKeysForModels(ctx context.Context, models []interface{}) (*keyrotation.RotationStats, error) {
// Create master context with timeout
masterCtx, cancel := context.WithTimeout(ctx, r.options.Timeout)
defer cancel()
// Track overall stats
overallStats := &keyrotation.RotationStats{
StartTime: time.Now(),
Errors: make([]string, 0),
}
// Create key rotator
rotator, err := keyrotation.NewKeyRotator(r.db, r.oldService, r.newService, r.auditor)
if err != nil {
return overallStats, fmt.Errorf("failed to create key rotator: %w", err)
}
// Apply options
rotator.SetDryRun(r.options.DryRun)
rotator.SetBatchSize(r.options.BatchSize)
rotator.SetMaxErrors(r.options.MaxErrors)
// Log the start of rotation
r.auditor.LogKeyRotationEventWithDescription(
"starting",
"pending",
true,
fmt.Sprintf("Starting key rotation for %d model types (dry run: %v)", len(models), r.options.DryRun),
0,
)
// Process all models (sequentially)
for _, model := range models {
// Check if context is canceled
select {
case <-masterCtx.Done():
overallStats.Errors = append(overallStats.Errors, fmt.Sprintf("key rotation aborted: %v", masterCtx.Err()))
return overallStats, masterCtx.Err()
default:
// Continue processing
}
// Get model type info
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Ptr {
modelType = modelType.Elem()
}
modelName := modelType.Name()
// Run pre-rotation hook if any
if err := r.runHook("pre_rotation_"+modelName, model); err != nil {
overallStats.Errors = append(overallStats.Errors, fmt.Sprintf("pre-rotation hook failed for %s: %v", modelName, err))
continue
}
// Log model rotation start
r.auditor.LogKeyRotationEventWithDescription(
"starting",
"pending",
true,
fmt.Sprintf("Starting key rotation for model: %s", modelName),
0,
)
// Create a worker context with timeout
workerCtx, workerCancel := context.WithTimeout(masterCtx, r.options.WorkerTimeout)
// Create a goroutine to handle timeouts
rotationDone := make(chan struct{})
var modelStats *keyrotation.RotationStats
var rotationErr error
go func() {
// Perform the actual rotation
modelStats, rotationErr = rotator.RotateKeys(model, "")
close(rotationDone)
}()
// Wait for rotation to complete or timeout
select {
case <-workerCtx.Done():
if workerCtx.Err() == context.DeadlineExceeded {
errorMsg := fmt.Sprintf("key rotation for model %s timed out after %v", modelName, r.options.WorkerTimeout)
overallStats.Errors = append(overallStats.Errors, errorMsg)
// Log timeout error
r.auditor.LogKeyRotationEventWithDescription(
"old",
"new",
false,
errorMsg,
0,
)
}
case <-rotationDone:
// Rotation completed
}
// Clean up the worker context
workerCancel()
// Check for rotation errors
if rotationErr != nil {
overallStats.Errors = append(overallStats.Errors, fmt.Sprintf("failed to rotate keys for %s: %v", modelName, rotationErr))
// Log rotation error
r.auditor.LogKeyRotationEventWithDescription(
"old",
"new",
false,
fmt.Sprintf("Key rotation failed for model %s: %v", modelName, rotationErr),
0,
)
continue
}
// Update overall stats
if modelStats != nil {
overallStats.TotalRecords += modelStats.TotalRecords
overallStats.ProcessedRecords += modelStats.ProcessedRecords
overallStats.SkippedRecords += modelStats.SkippedRecords
overallStats.FailedRecords += modelStats.FailedRecords
overallStats.Errors = append(overallStats.Errors, modelStats.Errors...)
// Call progress callback if set
if r.options.ProgressCallback != nil {
r.options.ProgressCallback(modelName, modelStats.ProcessedRecords, modelStats.TotalRecords)
}
// Log progress
successRate := 0.0
if modelStats.TotalRecords > 0 {
successRate = float64(modelStats.ProcessedRecords) / float64(modelStats.TotalRecords) * 100
}
r.auditor.LogKeyRotationEventWithDescription(
"old",
"new",
true,
fmt.Sprintf("Completed key rotation for model %s: %d/%d records (%.1f%%) processed, %d skipped, %d failed",
modelName, modelStats.ProcessedRecords, modelStats.TotalRecords, successRate,
modelStats.SkippedRecords, modelStats.FailedRecords),
0,
)
}
// Run post-rotation hook if any
if err := r.runHook("post_rotation_"+modelName, model); err != nil {
overallStats.Errors = append(overallStats.Errors, fmt.Sprintf("post-rotation hook failed for %s: %v", modelName, err))
}
}
// Complete overall stats
overallStats.EndTime = time.Now()
overallStats.ElapsedTime = overallStats.EndTime.Sub(overallStats.StartTime)
// Calculate overall success rate
successRate := 0.0
if overallStats.TotalRecords > 0 {
successRate = float64(overallStats.ProcessedRecords) / float64(overallStats.TotalRecords) * 100
}
// Log completion
r.auditor.LogKeyRotationEventWithDescription(
"old",
"new",
len(overallStats.Errors) == 0,
fmt.Sprintf("Completed key rotation for all models: %d/%d records (%.1f%%) processed, %d skipped, %d failed, %d errors in %s",
overallStats.ProcessedRecords, overallStats.TotalRecords, successRate,
overallStats.SkippedRecords, overallStats.FailedRecords, len(overallStats.Errors),
overallStats.ElapsedTime),
0,
)
return overallStats, nil
}
// FindModelsWithEncryptedFields automatically finds all database models with encrypted fields
func (r *RotationUtility) FindModelsWithEncryptedFields() ([]interface{}, error) {
// This is a placeholder - in a real implementation, we would scan the codebase
// or database schema to automatically detect models with encrypted fields
// Since that requires knowledge of the codebase structure, this would be
// customized for the specific application
return []interface{}{}, fmt.Errorf("automatic model detection not implemented, provide models explicitly")
}
// ValidateRotation tests the key rotation on sample records without saving changes
func (r *RotationUtility) ValidateRotation(models []interface{}) (map[string]bool, error) {
results := make(map[string]bool)
// Save current options to restore later
originalDryRun := r.options.DryRun
originalBatchSize := r.options.BatchSize
// Set temporary options for validation
r.options.DryRun = true
r.options.BatchSize = 10 // Test with small batch
// Create a context with short timeout
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
defer cancel()
// Run rotation with dry run mode
stats, err := r.RotateKeysForModels(ctx, models)
// Restore original options
r.options.DryRun = originalDryRun
r.options.BatchSize = originalBatchSize
if err != nil {
return results, fmt.Errorf("validation failed: %w", err)
}
// Process results for each model
for _, model := range models {
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Ptr {
modelType = modelType.Elem()
}
modelName := modelType.Name()
// Check if there were errors for this model
hasModelErrors := false
for _, errMsg := range stats.Errors {
if strings.Contains(errMsg, modelName) {
hasModelErrors = true
break
}
}
results[modelName] = !hasModelErrors
}
return results, nil
}
// CreateEncryptionMigrationPlan creates a detailed plan for migrating data to a new encryption key
func (r *RotationUtility) CreateEncryptionMigrationPlan(models []interface{}) (*EncryptionMigrationPlan, error) {
plan := &EncryptionMigrationPlan{
ModelPlans: make(map[string]*ModelMigrationPlan),
EstimatedDuration: 0,
EstimatedRecords: 0,
RecommendedOptions: r.options, // Start with current options
}
// Calculate record counts for each model
totalRecords := 0
for _, model := range models {
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Ptr {
modelType = modelType.Elem()
}
modelName := modelType.Name()
// Get record count
var count int64
if err := r.db.Model(model).Count(&count).Error; err != nil {
return nil, fmt.Errorf("failed to count records for %s: %w", modelName, err)
}
encryptedFields := r.identifyEncryptedFields(model)
// Create model plan
modelPlan := &ModelMigrationPlan{
ModelName: modelName,
RecordCount: int(count),
EstimatedTime: r.estimateMigrationTime(int(count), len(encryptedFields)),
EncryptedFields: encryptedFields,
BatchSizeRec: r.calculateOptimalBatchSize(int(count)),
}
plan.ModelPlans[modelName] = modelPlan
totalRecords += int(count)
plan.EstimatedDuration += modelPlan.EstimatedTime
}
plan.EstimatedRecords = totalRecords
// Calculate optimal batch size and parallelism based on total record count
plan.RecommendedOptions.BatchSize = r.calculateOptimalBatchSize(totalRecords)
plan.RecommendedOptions.Parallelism = r.calculateOptimalParallelism(totalRecords)
return plan, nil
}
// identifyEncryptedFields finds all encrypted fields in a model
func (r *RotationUtility) identifyEncryptedFields(model interface{}) []string {
fields := []string{}
// Get model value and type
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Ptr {
modelType = modelType.Elem()
}
// Skip if not a struct
if modelType.Kind() != reflect.Struct {
return fields
}
// Scan all fields for encrypted ones
for i := 0; i < modelType.NumField(); i++ {
field := modelType.Field(i)
// Look for fields starting with "Encrypted"
if strings.HasPrefix(field.Name, "Encrypted") && field.Type.Kind() == reflect.String {
fields = append(fields, field.Name)
}
}
return fields
}
// calculateOptimalBatchSize determines the optimal batch size based on record count
func (r *RotationUtility) calculateOptimalBatchSize(recordCount int) int {
// This is a simplistic approach - in a real system, this would be based on
// benchmarking and system characteristics
if recordCount < 1000 {
return 100
} else if recordCount < 10000 {
return 250
} else if recordCount < 100000 {
return 500
} else {
return 1000
}
}
// calculateOptimalParallelism determines the optimal parallelism level
func (r *RotationUtility) calculateOptimalParallelism(recordCount int) int {
// Simple heuristic - adjust based on actual system performance
cpuCount := runtime.NumCPU()
if recordCount < 10000 {
return 1
} else if recordCount < 100000 {
return min(2, cpuCount)
} else {
return min(4, cpuCount)
}
}
// min returns the minimum of two integers
func min(a, b int) int {
if a < b {
return a
}
return b
}
// estimateMigrationTime provides a rough estimate of time needed for migration
func (r *RotationUtility) estimateMigrationTime(recordCount, fieldCount int) time.Duration {
// This is a very rough estimate - in a real system, this would be based on
// benchmarking results and system characteristics
// Assume roughly 10ms per record per field
msPerRecordField := 10
// Calculate total time in milliseconds
totalTimeMs := recordCount * fieldCount * msPerRecordField
// Add overhead
totalTimeMs = int(float64(totalTimeMs) * 1.2) // 20% overhead
return time.Duration(totalTimeMs) * time.Millisecond
}
// EncryptionMigrationPlan contains the complete plan for migration
type EncryptionMigrationPlan struct {
ModelPlans map[string]*ModelMigrationPlan `json:"model_plans"`
EstimatedDuration time.Duration `json:"estimated_duration"`
EstimatedRecords int `json:"estimated_records"`
RecommendedOptions RotationOptions `json:"recommended_options"`
}
// ModelMigrationPlan contains migration details for a specific model
type ModelMigrationPlan struct {
ModelName string `json:"model_name"`
RecordCount int `json:"record_count"`
EstimatedTime time.Duration `json:"estimated_time"`
EncryptedFields []string `json:"encrypted_fields"`
BatchSizeRec int `json:"batch_size_recommendation"`
}
+271
View File
@@ -0,0 +1,271 @@
package audit
import (
"encoding/json"
"fmt"
"io"
"os"
"strings"
"sync"
"time"
)
// SecurityMonitor provides aggregate monitoring, alerting, and reporting for security events
type SecurityMonitor struct {
auditor *SecurityAuditor
statsMutex sync.RWMutex
eventCounts map[EventType]int
errorCounts map[string]int
lastEventTime map[EventType]time.Time
alertThresholds map[EventType]int
alertHandler AlertHandler
}
// AlertLevel represents the severity of a security alert
type AlertLevel string
// Alert levels
const (
AlertLevelInfo AlertLevel = "info"
AlertLevelWarning AlertLevel = "warning"
AlertLevelCritical AlertLevel = "critical"
)
// SecurityAlert represents a security alert to be sent to handlers
type SecurityAlert struct {
Timestamp time.Time
Level AlertLevel
EventType EventType
Message string
Count int
Details map[string]interface{}
}
// AlertHandler is the interface for handling security alerts
type AlertHandler interface {
HandleAlert(alert SecurityAlert)
}
// DefaultAlertHandler is a basic implementation of AlertHandler that logs to a file
type DefaultAlertHandler struct {
logFile string
writer io.Writer
writerLock sync.Mutex
}
// NewDefaultAlertHandler creates a new default alert handler
func NewDefaultAlertHandler(logFile string) (*DefaultAlertHandler, error) {
var writer io.Writer
if logFile == "" {
writer = os.Stdout
} else {
file, err := os.OpenFile(logFile, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
if err != nil {
return nil, fmt.Errorf("failed to open alert log file: %w", err)
}
writer = file
}
return &DefaultAlertHandler{
logFile: logFile,
writer: writer,
}, nil
}
// HandleAlert logs the alert to the configured output
func (h *DefaultAlertHandler) HandleAlert(alert SecurityAlert) {
h.writerLock.Lock()
defer h.writerLock.Unlock()
jsonData, err := json.Marshal(alert)
if err != nil {
fmt.Fprintf(h.writer, "Error marshaling alert: %v\n", err)
return
}
fmt.Fprintln(h.writer, string(jsonData))
}
// Close closes any open resources
func (h *DefaultAlertHandler) Close() error {
if h.logFile != "" {
if closer, ok := h.writer.(io.Closer); ok {
return closer.Close()
}
}
return nil
}
// NewSecurityMonitor creates a new SecurityMonitor
func NewSecurityMonitor(auditor *SecurityAuditor) *SecurityMonitor {
// Use provided auditor or global one if nil
if auditor == nil {
auditor = GetGlobalAuditor()
}
defaultHandler, _ := NewDefaultAlertHandler("")
return &SecurityMonitor{
auditor: auditor,
eventCounts: make(map[EventType]int),
errorCounts: make(map[string]int),
lastEventTime: make(map[EventType]time.Time),
alertThresholds: make(map[EventType]int),
alertHandler: defaultHandler,
}
}
// SetAlertHandler sets a custom alert handler
func (m *SecurityMonitor) SetAlertHandler(handler AlertHandler) {
m.alertHandler = handler
}
// SetAlertThreshold sets the threshold for when to generate alerts for a specific event type
func (m *SecurityMonitor) SetAlertThreshold(eventType EventType, threshold int) {
m.statsMutex.Lock()
defer m.statsMutex.Unlock()
m.alertThresholds[eventType] = threshold
}
// ProcessEvent processes a security event for monitoring
func (m *SecurityMonitor) ProcessEvent(event AuditEvent) {
m.statsMutex.Lock()
defer m.statsMutex.Unlock()
// Update event statistics
m.eventCounts[event.EventType]++
m.lastEventTime[event.EventType] = event.Timestamp
// Track errors
if !event.Success && event.Error != "" {
errorType := classifyError(event.Error)
m.errorCounts[errorType]++
// Alert on specific error types
if strings.Contains(event.Error, "unauthorized") ||
strings.Contains(event.Error, "permission") ||
strings.Contains(event.Error, "access denied") {
m.generateAlert(AlertLevelCritical, event.EventType,
fmt.Sprintf("Possible security breach detected: %s", event.Error),
map[string]interface{}{
"operation": event.Operation,
"error": event.Error,
"keyVersion": event.KeyVersion,
"modelType": event.ModelType,
})
}
}
// Check thresholds for alerting
threshold, hasThreshold := m.alertThresholds[event.EventType]
if hasThreshold && m.eventCounts[event.EventType] >= threshold {
if event.EventType == EventDecryptionFailure || event.EventType == EventEncryptionFailure {
m.generateAlert(AlertLevelWarning, event.EventType,
fmt.Sprintf("High number of %s events detected (%d)", event.EventType, m.eventCounts[event.EventType]),
map[string]interface{}{
"count": m.eventCounts[event.EventType],
"threshold": threshold,
})
} else if event.EventType == EventKeyRotation {
m.generateAlert(AlertLevelInfo, event.EventType,
fmt.Sprintf("Key rotation threshold reached (%d operations)", m.eventCounts[event.EventType]),
map[string]interface{}{
"count": m.eventCounts[event.EventType],
"threshold": threshold,
})
}
// Reset counter after alerting
m.eventCounts[event.EventType] = 0
}
}
// GenerateReport generates a report of security events for a time period
func (m *SecurityMonitor) GenerateReport(startTime, endTime time.Time, writer io.Writer) error {
m.statsMutex.RLock()
defer m.statsMutex.RUnlock()
report := struct {
TimeRange struct {
Start time.Time `json:"start"`
End time.Time `json:"end"`
} `json:"time_range"`
EventCounts map[EventType]int `json:"event_counts"`
ErrorCounts map[string]int `json:"error_counts"`
LastEventTimes map[EventType]time.Time `json:"last_event_times"`
GeneratedAt time.Time `json:"generated_at"`
}{
TimeRange: struct {
Start time.Time `json:"start"`
End time.Time `json:"end"`
}{
Start: startTime,
End: endTime,
},
EventCounts: m.eventCounts,
ErrorCounts: m.errorCounts,
LastEventTimes: m.lastEventTime,
GeneratedAt: time.Now(),
}
jsonData, err := json.MarshalIndent(report, "", " ")
if err != nil {
return fmt.Errorf("failed to marshal report: %w", err)
}
_, err = writer.Write(jsonData)
return err
}
// generateAlert creates and sends a security alert
func (m *SecurityMonitor) generateAlert(level AlertLevel, eventType EventType, message string, details map[string]interface{}) {
if m.alertHandler == nil {
return
}
alert := SecurityAlert{
Timestamp: time.Now(),
Level: level,
EventType: eventType,
Message: message,
Count: m.eventCounts[eventType],
Details: details,
}
go m.alertHandler.HandleAlert(alert)
}
// classifyError examines an error string and categorizes it
func classifyError(errorStr string) string {
errorStr = strings.ToLower(errorStr)
if strings.Contains(errorStr, "decrypt") {
return "decryption_error"
} else if strings.Contains(errorStr, "encrypt") {
return "encryption_error"
} else if strings.Contains(errorStr, "key") {
return "key_error"
} else if strings.Contains(errorStr, "permission") || strings.Contains(errorStr, "unauthorized") {
return "permission_error"
} else {
return "other_error"
}
}
// AttachToAuditor creates a wrapper function for the auditor's LogEvent method
// that processes events through the monitor before passing them to the original function.
// Returns the wrapped function that should be set on the auditor.
func (m *SecurityMonitor) AttachToAuditor() func(AuditEvent) {
originalLogEvent := m.auditor.LogEvent
// Create a wrapper function that processes events and then calls the original
return func(event AuditEvent) {
// Process the event for monitoring
m.ProcessEvent(event)
// Call the original LogEvent function
originalLogEvent(event)
}
}
+99
View File
@@ -0,0 +1,99 @@
package audit
import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
)
// MockAuditor is a mock implementation of an auditor
type MockAuditor struct {
mock.Mock
}
// LogEvent implements the required interface method
func (m *MockAuditor) LogEvent(event AuditEvent) {
m.Called(event)
}
// MockAlertHandler is a mock implementation of an AlertHandler
type MockAlertHandler struct {
mock.Mock
}
// HandleAlert implements the AlertHandler interface
func (m *MockAlertHandler) HandleAlert(alert SecurityAlert) {
m.Called(alert)
}
func TestSecurityMonitor(t *testing.T) {
// Create mocks
mockAuditor := new(MockAuditor)
mockAlertHandler := new(MockAlertHandler)
// Create the monitor
monitor := NewSecurityMonitor(mockAuditor)
monitor.SetAlertHandler(mockAlertHandler)
// Set up expectations
testEvent := AuditEvent{
Type: "key_rotation",
Description: "Key rotation completed",
Timestamp: time.Now(),
}
// The original auditor will be called
mockAuditor.On("LogEvent", testEvent).Return()
// Replace the auditor's LogEvent with our wrapped version
wrappedLogEvent := monitor.AttachToAuditor()
// Call the wrapped function
wrappedLogEvent(testEvent)
// Verify the expectations
mockAuditor.AssertExpectations(t)
// Test alert generation and handling
mockAlertHandler.On("HandleAlert", mock.Anything).Return()
errorEvent := AuditEvent{
Type: "error",
Description: "Failed to decrypt data: invalid key",
Timestamp: time.Now(),
Success: false,
}
// Process the error event directly to test alert generation
monitor.ProcessEvent(errorEvent)
// Verify alert was handled
mockAlertHandler.AssertExpectations(t)
// Test reporting functionality
report := monitor.GenerateReport()
assert.Contains(t, report.EventCounts, "key_rotation")
assert.Contains(t, report.ErrorCategories, "decryption_error")
}
func TestClassifyError(t *testing.T) {
testCases := []struct {
errorMsg string
expectedClass string
}{
{"failed to decrypt data", "decryption_error"},
{"encryption operation failed", "encryption_error"},
{"invalid key format", "key_error"},
{"unauthorized access to encryption key", "permission_error"},
{"some other random error", "other_error"},
}
for _, tc := range testCases {
t.Run(tc.errorMsg, func(t *testing.T) {
result := classifyError(tc.errorMsg)
assert.Equal(t, tc.expectedClass, result)
})
}
}
@@ -0,0 +1,563 @@
package audit
import (
"bytes"
"context"
"fmt"
"io"
"os"
"reflect"
"runtime"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/starfleetcptn/gomft/internal/encryption"
)
// TestingLevel represents the thoroughness of security tests
type TestingLevel int
const (
// BasicTesting includes essential encryption/decryption and key management tests
BasicTesting TestingLevel = iota
// ExtendedTesting adds key rotation, performance, and some edge cases
ExtendedTesting
// ComprehensiveTesting includes all tests plus stress tests, fuzzing, and security audit
ComprehensiveTesting
)
// TestSecretKey is a constant test key for testing purposes only
// Never use this in production
var TestSecretKey = []byte("01234567890123456789012345678901") // 32-byte key for AES-256
// SecurityTestingFramework provides comprehensive testing and benchmarking for the encryption system
type SecurityTestingFramework struct {
auditor *SecurityAuditor
monitor *SecurityMonitor
testOutputDir string
testLevel TestingLevel
logOutput io.Writer
verbose bool
mutex sync.Mutex
}
// TestResult represents the outcome of a security test
type TestResult struct {
Name string `json:"name"`
Success bool `json:"success"`
ElapsedTime time.Duration `json:"elapsed_time"`
Error string `json:"error,omitempty"`
Details string `json:"details,omitempty"`
}
// PerformanceMetrics contains performance data for encryption operations
type PerformanceMetrics struct {
OperationsPerSecond float64 `json:"operations_per_second"`
AverageLatency time.Duration `json:"average_latency"`
P95Latency time.Duration `json:"p95_latency"`
P99Latency time.Duration `json:"p99_latency"`
MemoryUsageMB float64 `json:"memory_usage_mb"`
CPUUsagePercent float64 `json:"cpu_usage_percent"`
}
// NewSecurityTestingFramework creates a new security testing framework
func NewSecurityTestingFramework(auditor *SecurityAuditor, monitor *SecurityMonitor) *SecurityTestingFramework {
if auditor == nil {
auditor = GetGlobalAuditor()
}
if monitor == nil {
monitor = NewSecurityMonitor(auditor)
}
return &SecurityTestingFramework{
auditor: auditor,
monitor: monitor,
testOutputDir: "security_test_results",
testLevel: BasicTesting,
logOutput: os.Stdout,
verbose: false,
}
}
// SetOutputDirectory sets the directory for test outputs
func (f *SecurityTestingFramework) SetOutputDirectory(dir string) {
f.mutex.Lock()
defer f.mutex.Unlock()
f.testOutputDir = dir
}
// SetTestingLevel sets the testing thoroughness level
func (f *SecurityTestingFramework) SetTestingLevel(level TestingLevel) {
f.mutex.Lock()
defer f.mutex.Unlock()
f.testLevel = level
}
// SetVerbose enables or disables verbose logging
func (f *SecurityTestingFramework) SetVerbose(verbose bool) {
f.mutex.Lock()
defer f.mutex.Unlock()
f.verbose = verbose
}
// SetLogOutput sets the output writer for test logs
func (f *SecurityTestingFramework) SetLogOutput(w io.Writer) {
f.mutex.Lock()
defer f.mutex.Unlock()
f.logOutput = w
}
// logf logs a message if verbose mode is enabled
func (f *SecurityTestingFramework) logf(format string, args ...interface{}) {
if f.verbose && f.logOutput != nil {
fmt.Fprintf(f.logOutput, format+"\n", args...)
}
}
// BenchmarkEncryptionPerformance measures the performance of encryption operations
func (f *SecurityTestingFramework) BenchmarkEncryptionPerformance(
service *encryption.EncryptionService,
dataSize int,
duration time.Duration,
) (*PerformanceMetrics, error) {
if service == nil {
return nil, fmt.Errorf("encryption service cannot be nil")
}
f.logf("Starting encryption performance benchmark (data size: %d bytes, duration: %s)", dataSize, duration)
// Generate test data
testData := make([]byte, dataSize)
for i := range testData {
testData[i] = byte(i % 256)
}
// Setup variables for benchmark
var (
operationCount uint64
totalLatency uint64
latencies []time.Duration
memStatsBefore runtime.MemStats
memStatsAfter runtime.MemStats
)
// Collect memory stats before
runtime.ReadMemStats(&memStatsBefore)
// Create context with timeout
ctx, cancel := context.WithTimeout(context.Background(), duration)
defer cancel()
// Record start time
startTime := time.Now()
// Run benchmark operations
var wg sync.WaitGroup
for i := 0; i < runtime.NumCPU(); i++ {
wg.Add(1)
go func() {
defer wg.Done()
localLatencies := make([]time.Duration, 0, 1000)
localData := make([]byte, len(testData))
copy(localData, testData)
for {
select {
case <-ctx.Done():
// Add local latencies to global latencies with lock
f.mutex.Lock()
latencies = append(latencies, localLatencies...)
f.mutex.Unlock()
return
default:
// Perform encrypt+decrypt operation and measure latency
opStart := time.Now()
// Encrypt
encrypted, err := service.Encrypt(localData)
if err != nil {
f.logf("Encryption error during benchmark: %v", err)
continue
}
// Decrypt
_, err = service.Decrypt(encrypted)
if err != nil {
f.logf("Decryption error during benchmark: %v", err)
continue
}
// Record latency
latency := time.Since(opStart)
localLatencies = append(localLatencies, latency)
// Update metrics
atomic.AddUint64(&operationCount, 1)
atomic.AddUint64(&totalLatency, uint64(latency))
}
}
}()
}
// Wait for the benchmark to complete
wg.Wait()
// Record end time
endTime := time.Now()
actualDuration := endTime.Sub(startTime)
// Collect memory stats after
runtime.ReadMemStats(&memStatsAfter)
// Calculate performance metrics
ops := atomic.LoadUint64(&operationCount)
if ops == 0 {
return nil, fmt.Errorf("no operations completed during benchmark")
}
// Sort latencies for percentile calculation
f.mutex.Lock()
latenciesLen := len(latencies)
f.mutex.Unlock()
// Calculate results
opsPerSec := float64(ops) / actualDuration.Seconds()
avgLatency := time.Duration(atomic.LoadUint64(&totalLatency) / ops)
// Calculate memory usage
memUsageMB := float64(memStatsAfter.Alloc-memStatsBefore.Alloc) / 1024 / 1024
// Calculate CPU usage (approximate based on operations)
cpuUsage := float64(ops) / float64(runtime.NumCPU()) / actualDuration.Seconds() * 100
if cpuUsage > 100 {
cpuUsage = 100
}
// Calculate P95 and P99 latencies
var p95Latency, p99Latency time.Duration
if latenciesLen > 0 {
f.mutex.Lock()
// Simple bubble sort for small sets (in production you'd use a more efficient sort)
for i := 0; i < latenciesLen; i++ {
for j := i + 1; j < latenciesLen; j++ {
if latencies[i] > latencies[j] {
latencies[i], latencies[j] = latencies[j], latencies[i]
}
}
}
p95Index := int(float64(latenciesLen) * 0.95)
p99Index := int(float64(latenciesLen) * 0.99)
if p95Index < latenciesLen {
p95Latency = latencies[p95Index]
}
if p99Index < latenciesLen {
p99Latency = latencies[p99Index]
}
f.mutex.Unlock()
}
metrics := &PerformanceMetrics{
OperationsPerSecond: opsPerSec,
AverageLatency: avgLatency,
P95Latency: p95Latency,
P99Latency: p99Latency,
MemoryUsageMB: memUsageMB,
CPUUsagePercent: cpuUsage,
}
f.logf("Encryption performance benchmark completed: %.2f ops/sec, avg latency: %s",
metrics.OperationsPerSecond, metrics.AverageLatency)
return metrics, nil
}
// VerifyKeyRotation tests the key rotation process
func (f *SecurityTestingFramework) VerifyKeyRotation(
oldService, newService *encryption.EncryptionService,
testData []byte,
) (*TestResult, error) {
startTime := time.Now()
result := &TestResult{
Name: "KeyRotationVerification",
}
if oldService == nil || newService == nil {
result.Success = false
result.Error = "encryption services cannot be nil"
return result, fmt.Errorf(result.Error)
}
f.logf("Verifying key rotation with %d bytes of test data", len(testData))
// Step 1: Encrypt with old key
encrypted, err := oldService.Encrypt(testData)
if err != nil {
result.Success = false
result.Error = fmt.Sprintf("failed to encrypt with old key: %v", err)
return result, fmt.Errorf(result.Error)
}
// Step 2: Verify old key can decrypt
decrypted, err := oldService.Decrypt(encrypted)
if err != nil {
result.Success = false
result.Error = fmt.Sprintf("failed to decrypt with old key: %v", err)
return result, fmt.Errorf(result.Error)
}
if string(decrypted) != string(testData) {
result.Success = false
result.Error = "decryption with old key produced different data"
return result, fmt.Errorf(result.Error)
}
// Step 3: Re-encrypt with new key
rotatedEncrypted, err := newService.Encrypt(decrypted)
if err != nil {
result.Success = false
result.Error = fmt.Sprintf("failed to re-encrypt with new key: %v", err)
return result, fmt.Errorf(result.Error)
}
// Step 4: Verify new key can decrypt
finalDecrypted, err := newService.Decrypt(rotatedEncrypted)
if err != nil {
result.Success = false
result.Error = fmt.Sprintf("failed to decrypt with new key: %v", err)
return result, fmt.Errorf(result.Error)
}
if string(finalDecrypted) != string(testData) {
result.Success = false
result.Error = "final decryption produced different data"
return result, fmt.Errorf(result.Error)
}
// Step 5: Verify new key cannot decrypt old data (different IV/salt)
_, err = newService.Decrypt(encrypted)
if err == nil {
result.Success = false
result.Error = "new key should not be able to decrypt data encrypted with old key"
return result, fmt.Errorf(result.Error)
}
result.Success = true
result.ElapsedTime = time.Since(startTime)
result.Details = fmt.Sprintf("Successfully verified key rotation process in %s", result.ElapsedTime)
f.logf("Key rotation verification successful")
return result, nil
}
// VerifyNoSensitiveDataInLogs checks that sensitive data is not exposed in logs
func (f *SecurityTestingFramework) VerifyNoSensitiveDataInLogs(sensitiveData string) (*TestResult, error) {
startTime := time.Now()
result := &TestResult{
Name: "SensitiveDataExposureCheck",
}
f.logf("Verifying sensitive data is not exposed in logs")
// Create test buffer for logs
logBuffer := new(logger)
// Create a temporary auditor that logs to our buffer
tempAuditor, err := New()
if err != nil {
result.Success = false
result.Error = fmt.Sprintf("failed to create test auditor: %v", err)
return result, fmt.Errorf(result.Error)
}
// Set log writer to our buffer
auditValue := reflect.ValueOf(tempAuditor).Elem()
if logField := auditValue.FieldByName("logWriter"); logField.IsValid() && logField.CanSet() {
logField.Set(reflect.ValueOf(logBuffer))
}
if errorField := auditValue.FieldByName("errorWriter"); errorField.IsValid() && errorField.CanSet() {
errorField.Set(reflect.ValueOf(logBuffer))
}
// Create a temporary encryption service for testing
os.Setenv("TEST_KEY", "dGVzdGtleXRlc3RrZXl0ZXN0a2V5dGVzdGtleXRlc3Q=") // base64 test key
keyManager := encryption.NewKeyManager("TEST_KEY")
err = keyManager.Initialize()
if err != nil {
result.Success = false
result.Error = fmt.Sprintf("failed to initialize key manager: %v", err)
return result, fmt.Errorf(result.Error)
}
encryptionService, err := encryption.NewEncryptionService(keyManager)
if err != nil {
result.Success = false
result.Error = fmt.Sprintf("failed to create encryption service: %v", err)
return result, fmt.Errorf(result.Error)
}
// Perform operations that should log
encryptedData, err := encryptionService.EncryptString(sensitiveData)
if err != nil {
result.Success = false
result.Error = fmt.Sprintf("failed to encrypt test data: %v", err)
return result, fmt.Errorf(result.Error)
}
// Log various events with the sensitive data
tempAuditor.LogEncryptionEvent("test_encrypt", "password", "TestModel", true, nil, "v1", 0, time.Millisecond)
tempAuditor.LogDecryptionEvent("test_decrypt", "password", "TestModel", true, nil, "v1", 0, time.Millisecond)
tempAuditor.LogKeyRotationEvent("v1", "v2", true, nil, 0)
// Force an error log that might contain sensitive data
tempAuditor.LogDecryptionEvent("test_error", "password", "TestModel", false,
fmt.Errorf("failed to decrypt: %s", sensitiveData), "v1", 0, time.Millisecond)
// Get the log contents
logContents := logBuffer.String()
// Check if the sensitive data appears in the logs
if strings.Contains(logContents, sensitiveData) {
result.Success = false
result.Error = "sensitive data was found in the logs"
return result, fmt.Errorf(result.Error)
}
// Also check for the encrypted version
if strings.Contains(logContents, encryptedData) {
result.Success = false
result.Error = "encrypted sensitive data was found in the logs"
return result, fmt.Errorf(result.Error)
}
result.Success = true
result.ElapsedTime = time.Since(startTime)
result.Details = "Successfully verified that sensitive data is properly sanitized in logs"
f.logf("Sensitive data exposure check passed")
return result, nil
}
// Custom logger for testing
type logger struct {
buffer bytes.Buffer
mu sync.Mutex
}
func (l *logger) Write(p []byte) (n int, err error) {
l.mu.Lock()
defer l.mu.Unlock()
return l.buffer.Write(p)
}
func (l *logger) String() string {
l.mu.Lock()
defer l.mu.Unlock()
return l.buffer.String()
}
// RunAllTests executes all security tests based on the configured test level
func (f *SecurityTestingFramework) RunAllTests(encryptionService *encryption.EncryptionService) ([]*TestResult, error) {
results := make([]*TestResult, 0)
// Basic tests
basicTests := []func(*encryption.EncryptionService) (*TestResult, error){
f.testEncryptionDecryption,
f.testEmptyData,
f.testLargeData,
}
// Extended tests
extendedTests := []func(*encryption.EncryptionService) (*TestResult, error){
f.testPerformance,
f.testConcurrentAccess,
f.testKeyVersioning,
}
// Comprehensive tests
comprehensiveTests := []func(*encryption.EncryptionService) (*TestResult, error){
f.testFuzzedInput,
f.testKeyRotation,
f.testErrorHandling,
f.testSensitiveDataExposure,
}
// Run basic tests
for _, test := range basicTests {
result, err := test(encryptionService)
if err != nil {
f.logf("Test %s failed: %v", result.Name, err)
}
results = append(results, result)
}
// Run extended tests if level is high enough
if f.testLevel >= ExtendedTesting {
for _, test := range extendedTests {
result, err := test(encryptionService)
if err != nil {
f.logf("Test %s failed: %v", result.Name, err)
}
results = append(results, result)
}
}
// Run comprehensive tests if level is highest
if f.testLevel >= ComprehensiveTesting {
for _, test := range comprehensiveTests {
result, err := test(encryptionService)
if err != nil {
f.logf("Test %s failed: %v", result.Name, err)
}
results = append(results, result)
}
}
return results, nil
}
// Test implementations (placeholders - these would be implemented with real tests)
func (f *SecurityTestingFramework) testEncryptionDecryption(s *encryption.EncryptionService) (*TestResult, error) {
// This is a placeholder - in a real implementation, this would perform actual tests
return &TestResult{Name: "EncryptionDecryption", Success: true}, nil
}
func (f *SecurityTestingFramework) testEmptyData(s *encryption.EncryptionService) (*TestResult, error) {
return &TestResult{Name: "EmptyData", Success: true}, nil
}
func (f *SecurityTestingFramework) testLargeData(s *encryption.EncryptionService) (*TestResult, error) {
return &TestResult{Name: "LargeData", Success: true}, nil
}
func (f *SecurityTestingFramework) testPerformance(s *encryption.EncryptionService) (*TestResult, error) {
return &TestResult{Name: "Performance", Success: true}, nil
}
func (f *SecurityTestingFramework) testConcurrentAccess(s *encryption.EncryptionService) (*TestResult, error) {
return &TestResult{Name: "ConcurrentAccess", Success: true}, nil
}
func (f *SecurityTestingFramework) testKeyVersioning(s *encryption.EncryptionService) (*TestResult, error) {
return &TestResult{Name: "KeyVersioning", Success: true}, nil
}
func (f *SecurityTestingFramework) testFuzzedInput(s *encryption.EncryptionService) (*TestResult, error) {
return &TestResult{Name: "FuzzedInput", Success: true}, nil
}
func (f *SecurityTestingFramework) testKeyRotation(s *encryption.EncryptionService) (*TestResult, error) {
return &TestResult{Name: "KeyRotation", Success: true}, nil
}
func (f *SecurityTestingFramework) testErrorHandling(s *encryption.EncryptionService) (*TestResult, error) {
return &TestResult{Name: "ErrorHandling", Success: true}, nil
}
func (f *SecurityTestingFramework) testSensitiveDataExposure(s *encryption.EncryptionService) (*TestResult, error) {
return &TestResult{Name: "SensitiveDataExposure", Success: true}, nil
}
@@ -0,0 +1,196 @@
package audit
import (
"bytes"
"os"
"reflect"
"testing"
"time"
"github.com/starfleetcptn/gomft/internal/encryption"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func setupTestFramework(t *testing.T) (*SecurityTestingFramework, *bytes.Buffer) {
// Create audit log buffer
logBuffer := new(bytes.Buffer)
// Create auditor
auditor, err := New()
require.NoError(t, err)
// Set auditor to use buffer
auditValue := reflect.ValueOf(auditor).Elem()
if logField := auditValue.FieldByName("logWriter"); logField.IsValid() && logField.CanSet() {
logField.Set(reflect.ValueOf(logBuffer))
}
if errorField := auditValue.FieldByName("errorWriter"); errorField.IsValid() && errorField.CanSet() {
errorField.Set(reflect.ValueOf(logBuffer))
}
// Create monitor
monitor := NewSecurityMonitor(auditor)
// Create framework
framework := NewSecurityTestingFramework(auditor, monitor)
framework.SetVerbose(true)
return framework, logBuffer
}
func setupTestEncryptionService(t *testing.T) *encryption.EncryptionService {
// Setup test key
os.Setenv("TEST_ENCRYPTION_KEY", "dGVzdGtleXRlc3RrZXl0ZXN0a2V5dGVzdGtleXRlc3Q=") // base64 test key
t.Cleanup(func() {
os.Unsetenv("TEST_ENCRYPTION_KEY")
})
// Create key manager
keyManager := encryption.NewKeyManager("TEST_ENCRYPTION_KEY")
err := keyManager.Initialize()
require.NoError(t, err)
// Create encryption service
service, err := encryption.NewEncryptionService(keyManager)
require.NoError(t, err)
return service
}
func TestNewSecurityTestingFramework(t *testing.T) {
auditor, err := New()
require.NoError(t, err)
monitor := NewSecurityMonitor(auditor)
framework := NewSecurityTestingFramework(auditor, monitor)
assert.Equal(t, auditor, framework.auditor)
assert.Equal(t, monitor, framework.monitor)
assert.Equal(t, "security_test_results", framework.testOutputDir)
assert.Equal(t, BasicTesting, framework.testLevel)
assert.Equal(t, os.Stdout, framework.logOutput)
assert.False(t, framework.verbose)
}
func TestSecurityTestingFramework_SetMethods(t *testing.T) {
framework, _ := setupTestFramework(t)
// Test SetOutputDirectory
framework.SetOutputDirectory("test_dir")
assert.Equal(t, "test_dir", framework.testOutputDir)
// Test SetTestingLevel
framework.SetTestingLevel(ComprehensiveTesting)
assert.Equal(t, ComprehensiveTesting, framework.testLevel)
// Test SetVerbose
framework.SetVerbose(true)
assert.True(t, framework.verbose)
// Test SetLogOutput
buffer := new(bytes.Buffer)
framework.SetLogOutput(buffer)
assert.Equal(t, buffer, framework.logOutput)
}
func TestSecurityTestingFramework_BenchmarkEncryptionPerformance(t *testing.T) {
framework, _ := setupTestFramework(t)
service := setupTestEncryptionService(t)
// Run a very short benchmark
metrics, err := framework.BenchmarkEncryptionPerformance(service, 1024, 100*time.Millisecond)
require.NoError(t, err)
// Verify metrics are populated
assert.True(t, metrics.OperationsPerSecond > 0)
assert.True(t, metrics.AverageLatency > 0)
assert.True(t, metrics.MemoryUsageMB >= 0)
assert.True(t, metrics.CPUUsagePercent >= 0)
}
func TestSecurityTestingFramework_VerifyKeyRotation(t *testing.T) {
framework, _ := setupTestFramework(t)
// Setup two different encryption services with different keys
oldKeyEnv := "TEST_OLD_KEY"
newKeyEnv := "TEST_NEW_KEY"
os.Setenv(oldKeyEnv, "b2xka2V5b2xka2V5b2xka2V5b2xka2V5b2xka2V5b2xk")
os.Setenv(newKeyEnv, "bmV3a2V5bmV3a2V5bmV3a2V5bmV3a2V5bmV3a2V5bmV3")
t.Cleanup(func() {
os.Unsetenv(oldKeyEnv)
os.Unsetenv(newKeyEnv)
})
// Create old key manager and service
oldKeyManager := encryption.NewKeyManager(oldKeyEnv)
err := oldKeyManager.Initialize()
require.NoError(t, err)
oldService, err := encryption.NewEncryptionService(oldKeyManager)
require.NoError(t, err)
// Create new key manager and service
newKeyManager := encryption.NewKeyManager(newKeyEnv)
err = newKeyManager.Initialize()
require.NoError(t, err)
newService, err := encryption.NewEncryptionService(newKeyManager)
require.NoError(t, err)
// Test data
testData := []byte("This is some test data for key rotation verification")
// Run verification
result, err := framework.VerifyKeyRotation(oldService, newService, testData)
require.NoError(t, err)
assert.True(t, result.Success)
assert.Contains(t, result.Details, "Successfully verified key rotation")
}
func TestSecurityTestingFramework_VerifyNoSensitiveDataInLogs(t *testing.T) {
framework, _ := setupTestFramework(t)
// Sensitive data to check
sensitiveData := "very_sensitive_password_123!"
// Run verification
result, err := framework.VerifyNoSensitiveDataInLogs(sensitiveData)
require.NoError(t, err)
assert.True(t, result.Success)
assert.Contains(t, result.Details, "Successfully verified that sensitive data is properly sanitized")
}
func TestSecurityTestingFramework_RunAllTests(t *testing.T) {
framework, _ := setupTestFramework(t)
service := setupTestEncryptionService(t)
// Run tests at basic level
results, err := framework.RunAllTests(service)
require.NoError(t, err)
// Should have 3 basic tests
assert.Equal(t, 3, len(results))
// Set to extended level and run again
framework.SetTestingLevel(ExtendedTesting)
results, err = framework.RunAllTests(service)
require.NoError(t, err)
// Should have 3 basic + 3 extended tests
assert.Equal(t, 6, len(results))
// Set to comprehensive level and run again
framework.SetTestingLevel(ComprehensiveTesting)
results, err = framework.RunAllTests(service)
require.NoError(t, err)
// Should have 3 basic + 3 extended + 4 comprehensive tests
assert.Equal(t, 10, len(results))
}