diff --git a/cmd/auth-server/main.go b/cmd/auth-server/main.go index 5ae01ce..19c7ac9 100644 --- a/cmd/auth-server/main.go +++ b/cmd/auth-server/main.go @@ -88,6 +88,20 @@ func main() { userService := services.NewUserService(repos, logger) clientService := services.NewClientService(repos, logger) authService := services.NewAuthService(repos, cfg, logger) + roleService := services.NewRoleService( + repos.Role, + repos.Permission, + repos.UserRole, + repos.RolePermission, + auditService, + logger, + ) + permissionService := services.NewPermissionService( + repos.Permission, + repos.UserRole, + repos.RolePermission, + logger, + ) // Start background email queue processor go func() { @@ -107,6 +121,24 @@ func main() { } }() + // Start background token cleanup worker (runs every hour) + go func() { + ticker := time.NewTicker(1 * time.Hour) + defer ticker.Stop() + + for { + select { + case <-rootCtx.Done(): + logger.Info("Token cleanup worker shutting down") + return + case <-ticker.C: + if err := authService.CleanupExpiredTokens(rootCtx); err != nil { + logger.WithError(err).Error("Failed to cleanup expired tokens") + } + } + } + }() + // Setup Gin router if cfg.GinMode != "" { gin.SetMode(cfg.GinMode) @@ -127,9 +159,11 @@ func main() { userHandler := handlers.NewUserHandler(userService, logger) clientHandler := handlers.NewClientHandler(clientService, logger) oauthHandler := handlers.NewOAuthHandler(tenantService, userService, clientService, authService, logger) + rbacHandler := handlers.NewRBACHandler(roleService, permissionService, auditService, logger) + auditHandler := handlers.NewAuditHandler(auditService, logger) // Setup routes - setupRoutes(cfg, db, redisClient, router, tenantHandler, userHandler, clientHandler, oauthHandler) + setupRoutes(cfg, db, redisClient, router, tenantHandler, userHandler, clientHandler, oauthHandler, rbacHandler, auditHandler) // Create HTTP server server := &http.Server{ @@ -202,6 +236,8 @@ func setupRoutes( userHandler *handlers.UserHandler, clientHandler *handlers.ClientHandler, oauthHandler *handlers.OAuthHandler, + rbacHandler *handlers.RBACHandler, + auditHandler *handlers.AuditHandler, ) { // Health check endpoint router.GET("/health", func(c *gin.Context) { @@ -259,8 +295,14 @@ func setupRoutes( // Client management clientHandler.RegisterRoutes(api.Group("/clients")) + } + // RBAC and Audit routes register their own /v1/* prefixes and use + // RequireAuth middleware internally via the handler. + rbacHandler.RegisterRoutes(router.Group("")) + auditHandler.RegisterRoutes(router.Group("")) + // Serve static files and templates router.Static("/static", "./static") diff --git a/internal/database/database.go b/internal/database/database.go index eb87b6b..903087d 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -145,6 +145,183 @@ func Migrate(db *gorm.DB) error { `CREATE INDEX IF NOT EXISTS idx_refresh_tokens_client_id ON refresh_tokens(client_id)`, `CREATE INDEX IF NOT EXISTS idx_refresh_tokens_user_id ON refresh_tokens(user_id)`, `CREATE INDEX IF NOT EXISTS idx_refresh_tokens_expires_at ON refresh_tokens(expires_at)`, + + // ── Phase 1 Features ──────────────────────────────────────────────── + + // Enhanced user management columns + `ALTER TABLE users ADD COLUMN IF NOT EXISTS status VARCHAR(20) NOT NULL DEFAULT 'pending'`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS email_verified BOOLEAN NOT NULL DEFAULT FALSE`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS phone_number VARCHAR(50)`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS phone_verified BOOLEAN NOT NULL DEFAULT FALSE`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS first_name VARCHAR(255)`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS last_name VARCHAR(255)`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS avatar VARCHAR(500)`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS locale VARCHAR(10) DEFAULT 'en'`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS timezone VARCHAR(50) DEFAULT 'UTC'`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS last_login_at TIMESTAMP WITH TIME ZONE`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS last_login_ip VARCHAR(45)`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS failed_login_attempts INTEGER NOT NULL DEFAULT 0`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS locked_at TIMESTAMP WITH TIME ZONE`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS locked_until TIMESTAMP WITH TIME ZONE`, + `ALTER TABLE users ADD COLUMN IF NOT EXISTS metadata JSONB DEFAULT '{}'`, + `CREATE INDEX IF NOT EXISTS idx_users_status ON users(status)`, + `CREATE INDEX IF NOT EXISTS idx_users_email_verified ON users(email_verified)`, + + // Permissions table + `CREATE TABLE IF NOT EXISTS permissions ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + name VARCHAR(255) NOT NULL UNIQUE, + display_name VARCHAR(255) NOT NULL, + description TEXT, + resource VARCHAR(255) NOT NULL, + action VARCHAR(255) NOT NULL, + is_system BOOLEAN NOT NULL DEFAULT FALSE, + created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(), + updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(), + deleted_at TIMESTAMP WITH TIME ZONE + )`, + `CREATE INDEX IF NOT EXISTS idx_permissions_resource ON permissions(resource)`, + `CREATE INDEX IF NOT EXISTS idx_permissions_action ON permissions(action)`, + + // Roles table + `CREATE TABLE IF NOT EXISTS roles ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + tenant_id UUID NOT NULL, + name VARCHAR(255) NOT NULL, + display_name VARCHAR(255) NOT NULL, + description TEXT, + is_system BOOLEAN NOT NULL DEFAULT FALSE, + is_active BOOLEAN NOT NULL DEFAULT TRUE, + created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(), + updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(), + deleted_at TIMESTAMP WITH TIME ZONE + )`, + `CREATE UNIQUE INDEX IF NOT EXISTS idx_roles_tenant_name ON roles(tenant_id, name) WHERE deleted_at IS NULL`, + `CREATE INDEX IF NOT EXISTS idx_roles_tenant_id ON roles(tenant_id)`, + + // User roles junction table + `CREATE TABLE IF NOT EXISTS user_roles ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + tenant_id UUID NOT NULL, + user_id UUID NOT NULL, + role_id UUID NOT NULL, + granted_by UUID NOT NULL, + granted_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(), + expires_at TIMESTAMP WITH TIME ZONE + )`, + `CREATE UNIQUE INDEX IF NOT EXISTS idx_user_roles_user_role ON user_roles(tenant_id, user_id, role_id)`, + `CREATE INDEX IF NOT EXISTS idx_user_roles_user_id ON user_roles(user_id)`, + `CREATE INDEX IF NOT EXISTS idx_user_roles_role_id ON user_roles(role_id)`, + + // Role permissions junction table + `CREATE TABLE IF NOT EXISTS role_permissions ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + role_id UUID NOT NULL, + permission_id UUID NOT NULL, + granted_by UUID NOT NULL, + granted_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW() + )`, + `CREATE UNIQUE INDEX IF NOT EXISTS idx_role_permissions_role_perm ON role_permissions(role_id, permission_id)`, + `CREATE INDEX IF NOT EXISTS idx_role_permissions_role_id ON role_permissions(role_id)`, + `CREATE INDEX IF NOT EXISTS idx_role_permissions_permission_id ON role_permissions(permission_id)`, + + // Audit logs table + `CREATE TABLE IF NOT EXISTS audit_logs ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + tenant_id UUID NOT NULL, + user_id UUID, + client_id UUID, + action VARCHAR(255) NOT NULL, + resource VARCHAR(255) NOT NULL, + resource_id UUID, + ip_address VARCHAR(45), + user_agent TEXT, + request_id VARCHAR(255), + success BOOLEAN NOT NULL DEFAULT TRUE, + error_code VARCHAR(255), + error_message TEXT, + metadata JSONB DEFAULT '{}', + created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW() + )`, + `CREATE INDEX IF NOT EXISTS idx_audit_logs_tenant_id ON audit_logs(tenant_id)`, + `CREATE INDEX IF NOT EXISTS idx_audit_logs_user_id ON audit_logs(user_id)`, + `CREATE INDEX IF NOT EXISTS idx_audit_logs_action ON audit_logs(action)`, + `CREATE INDEX IF NOT EXISTS idx_audit_logs_resource ON audit_logs(resource)`, + `CREATE INDEX IF NOT EXISTS idx_audit_logs_created_at ON audit_logs(created_at)`, + `CREATE INDEX IF NOT EXISTS idx_audit_logs_success ON audit_logs(success)`, + + // Email templates table + `CREATE TABLE IF NOT EXISTS email_templates ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + tenant_id UUID NOT NULL, + name VARCHAR(255) NOT NULL, + subject VARCHAR(500) NOT NULL, + body_html TEXT, + body_text TEXT, + variables JSONB DEFAULT '[]', + is_active BOOLEAN NOT NULL DEFAULT TRUE, + created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(), + updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(), + deleted_at TIMESTAMP WITH TIME ZONE + )`, + `CREATE UNIQUE INDEX IF NOT EXISTS idx_email_templates_tenant_name ON email_templates(tenant_id, name) WHERE deleted_at IS NULL`, + `CREATE INDEX IF NOT EXISTS idx_email_templates_tenant_id ON email_templates(tenant_id)`, + + // Email queue table + `CREATE TABLE IF NOT EXISTS email_queue ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + tenant_id UUID NOT NULL, + user_id UUID, + to_email VARCHAR(255) NOT NULL, + to_name VARCHAR(255), + from_email VARCHAR(255) NOT NULL, + from_name VARCHAR(255), + subject VARCHAR(500) NOT NULL, + body_html TEXT, + body_text TEXT, + status VARCHAR(50) NOT NULL DEFAULT 'pending', + priority INTEGER NOT NULL DEFAULT 5, + attempts INTEGER NOT NULL DEFAULT 0, + max_attempts INTEGER NOT NULL DEFAULT 3, + last_error TEXT, + scheduled_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(), + sent_at TIMESTAMP WITH TIME ZONE, + created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(), + updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW() + )`, + `CREATE INDEX IF NOT EXISTS idx_email_queue_tenant_id ON email_queue(tenant_id)`, + `CREATE INDEX IF NOT EXISTS idx_email_queue_status ON email_queue(status)`, + `CREATE INDEX IF NOT EXISTS idx_email_queue_priority ON email_queue(priority)`, + `CREATE INDEX IF NOT EXISTS idx_email_queue_scheduled_at ON email_queue(scheduled_at)`, + + // Email verifications table + `CREATE TABLE IF NOT EXISTS email_verifications ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + tenant_id UUID NOT NULL, + user_id UUID NOT NULL, + email VARCHAR(255) NOT NULL, + code VARCHAR(255) NOT NULL UNIQUE, + expires_at TIMESTAMP WITH TIME ZONE NOT NULL, + verified_at TIMESTAMP WITH TIME ZONE, + created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW() + )`, + `CREATE INDEX IF NOT EXISTS idx_email_verifications_tenant_id ON email_verifications(tenant_id)`, + `CREATE INDEX IF NOT EXISTS idx_email_verifications_user_id ON email_verifications(user_id)`, + `CREATE INDEX IF NOT EXISTS idx_email_verifications_expires_at ON email_verifications(expires_at)`, + + // Password resets table + `CREATE TABLE IF NOT EXISTS password_resets ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + tenant_id UUID NOT NULL, + user_id UUID NOT NULL, + token VARCHAR(255) NOT NULL UNIQUE, + expires_at TIMESTAMP WITH TIME ZONE NOT NULL, + used_at TIMESTAMP WITH TIME ZONE, + created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW() + )`, + `CREATE INDEX IF NOT EXISTS idx_password_resets_tenant_id ON password_resets(tenant_id)`, + `CREATE INDEX IF NOT EXISTS idx_password_resets_user_id ON password_resets(user_id)`, + `CREATE INDEX IF NOT EXISTS idx_password_resets_expires_at ON password_resets(expires_at)`, } // Execute each migration diff --git a/internal/repo/gorm/audit_log.go b/internal/repo/gorm/audit_log.go new file mode 100644 index 0000000..68de852 --- /dev/null +++ b/internal/repo/gorm/audit_log.go @@ -0,0 +1,149 @@ +package gorm + +import ( + "context" + "errors" + "fmt" + "time" + + "shieldgate/internal/models" + "shieldgate/internal/repo" + + "github.com/google/uuid" + "gorm.io/gorm" +) + +type auditLogRepository struct { + db *gorm.DB +} + +// NewAuditLogRepository creates a new audit log repository +func NewAuditLogRepository(db *gorm.DB) repo.AuditLogRepository { + return &auditLogRepository{db: db} +} + +func (r *auditLogRepository) Create(ctx context.Context, auditLog *models.AuditLog) error { + if err := r.db.WithContext(ctx).Create(auditLog).Error; err != nil { + return fmt.Errorf("failed to create audit log: %w", err) + } + return nil +} + +func (r *auditLogRepository) GetByID(ctx context.Context, tenantID, auditID uuid.UUID) (*models.AuditLog, error) { + var auditLog models.AuditLog + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND id = ?", tenantID, auditID). + First(&auditLog).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrResourceNotFound + } + return nil, fmt.Errorf("failed to get audit log by ID: %w", err) + } + return &auditLog, nil +} + +func (r *auditLogRepository) Query(ctx context.Context, query *models.AuditLogQuery) ([]*models.AuditLog, int64, error) { + var auditLogs []*models.AuditLog + var total int64 + + db := r.db.WithContext(ctx).Model(&models.AuditLog{}).Where("tenant_id = ?", query.TenantID) + + if query.UserID != nil { + db = db.Where("user_id = ?", query.UserID) + } + if query.ClientID != nil { + db = db.Where("client_id = ?", query.ClientID) + } + if query.Action != nil { + db = db.Where("action = ?", query.Action) + } + if query.Resource != "" { + db = db.Where("resource = ?", query.Resource) + } + if query.ResourceID != nil { + db = db.Where("resource_id = ?", query.ResourceID) + } + if query.IPAddress != "" { + db = db.Where("ip_address = ?", query.IPAddress) + } + if query.Success != nil { + db = db.Where("success = ?", query.Success) + } + if query.StartDate != nil { + db = db.Where("created_at >= ?", query.StartDate) + } + if query.EndDate != nil { + db = db.Where("created_at <= ?", query.EndDate) + } + + if err := db.Count(&total).Error; err != nil { + return nil, 0, fmt.Errorf("failed to count audit logs: %w", err) + } + + limit := query.Limit + if limit <= 0 { + limit = 50 + } + + if err := db.Order("created_at DESC"). + Limit(limit). + Offset(query.Offset). + Find(&auditLogs).Error; err != nil { + return nil, 0, fmt.Errorf("failed to query audit logs: %w", err) + } + + return auditLogs, total, nil +} + +func (r *auditLogRepository) GetUserActivity(ctx context.Context, tenantID, userID uuid.UUID, limit, offset int) ([]*models.AuditLog, int64, error) { + var auditLogs []*models.AuditLog + var total int64 + + db := r.db.WithContext(ctx). + Model(&models.AuditLog{}). + Where("tenant_id = ? AND user_id = ?", tenantID, userID) + + if err := db.Count(&total).Error; err != nil { + return nil, 0, fmt.Errorf("failed to count user activity: %w", err) + } + + if err := db.Order("created_at DESC"). + Limit(limit). + Offset(offset). + Find(&auditLogs).Error; err != nil { + return nil, 0, fmt.Errorf("failed to get user activity: %w", err) + } + + return auditLogs, total, nil +} + +func (r *auditLogRepository) GetResourceActivity(ctx context.Context, tenantID uuid.UUID, resource string, resourceID uuid.UUID, limit, offset int) ([]*models.AuditLog, int64, error) { + var auditLogs []*models.AuditLog + var total int64 + + db := r.db.WithContext(ctx). + Model(&models.AuditLog{}). + Where("tenant_id = ? AND resource = ? AND resource_id = ?", tenantID, resource, resourceID) + + if err := db.Count(&total).Error; err != nil { + return nil, 0, fmt.Errorf("failed to count resource activity: %w", err) + } + + if err := db.Order("created_at DESC"). + Limit(limit). + Offset(offset). + Find(&auditLogs).Error; err != nil { + return nil, 0, fmt.Errorf("failed to get resource activity: %w", err) + } + + return auditLogs, total, nil +} + +func (r *auditLogRepository) DeleteOldLogs(ctx context.Context, olderThan time.Time) error { + if err := r.db.WithContext(ctx). + Where("created_at < ?", olderThan). + Delete(&models.AuditLog{}).Error; err != nil { + return fmt.Errorf("failed to delete old audit logs: %w", err) + } + return nil +} diff --git a/internal/repo/gorm/email.go b/internal/repo/gorm/email.go new file mode 100644 index 0000000..aa25f22 --- /dev/null +++ b/internal/repo/gorm/email.go @@ -0,0 +1,314 @@ +package gorm + +import ( + "context" + "errors" + "fmt" + "time" + + "shieldgate/internal/models" + "shieldgate/internal/repo" + + "github.com/google/uuid" + "gorm.io/gorm" +) + +// ─── EmailTemplate Repository ──────────────────────────────────────────────── + +type emailTemplateRepository struct { + db *gorm.DB +} + +// NewEmailTemplateRepository creates a new email template repository +func NewEmailTemplateRepository(db *gorm.DB) repo.EmailTemplateRepository { + return &emailTemplateRepository{db: db} +} + +func (r *emailTemplateRepository) Create(ctx context.Context, template *models.EmailTemplate) error { + if err := r.db.WithContext(ctx).Create(template).Error; err != nil { + return fmt.Errorf("failed to create email template: %w", err) + } + return nil +} + +func (r *emailTemplateRepository) GetByName(ctx context.Context, tenantID uuid.UUID, name string) (*models.EmailTemplate, error) { + var template models.EmailTemplate + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND name = ?", tenantID, name). + First(&template).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrResourceNotFound + } + return nil, fmt.Errorf("failed to get email template: %w", err) + } + return &template, nil +} + +func (r *emailTemplateRepository) Update(ctx context.Context, template *models.EmailTemplate) error { + if err := r.db.WithContext(ctx).Save(template).Error; err != nil { + return fmt.Errorf("failed to update email template: %w", err) + } + return nil +} + +func (r *emailTemplateRepository) Delete(ctx context.Context, tenantID uuid.UUID, name string) error { + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND name = ?", tenantID, name). + Delete(&models.EmailTemplate{}).Error; err != nil { + return fmt.Errorf("failed to delete email template: %w", err) + } + return nil +} + +func (r *emailTemplateRepository) List(ctx context.Context, tenantID uuid.UUID, limit, offset int) ([]*models.EmailTemplate, int64, error) { + var templates []*models.EmailTemplate + var total int64 + + if err := r.db.WithContext(ctx). + Model(&models.EmailTemplate{}). + Where("tenant_id = ?", tenantID). + Count(&total).Error; err != nil { + return nil, 0, fmt.Errorf("failed to count email templates: %w", err) + } + + if err := r.db.WithContext(ctx). + Where("tenant_id = ?", tenantID). + Order("name ASC"). + Limit(limit). + Offset(offset). + Find(&templates).Error; err != nil { + return nil, 0, fmt.Errorf("failed to list email templates: %w", err) + } + + return templates, total, nil +} + +// ─── EmailQueue Repository ──────────────────────────────────────────────────── + +type emailQueueRepository struct { + db *gorm.DB +} + +// NewEmailQueueRepository creates a new email queue repository +func NewEmailQueueRepository(db *gorm.DB) repo.EmailQueueRepository { + return &emailQueueRepository{db: db} +} + +func (r *emailQueueRepository) Create(ctx context.Context, email *models.EmailQueue) error { + if err := r.db.WithContext(ctx).Create(email).Error; err != nil { + return fmt.Errorf("failed to create email queue entry: %w", err) + } + return nil +} + +func (r *emailQueueRepository) GetByID(ctx context.Context, tenantID, emailID uuid.UUID) (*models.EmailQueue, error) { + var email models.EmailQueue + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND id = ?", tenantID, emailID). + First(&email).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrResourceNotFound + } + return nil, fmt.Errorf("failed to get email queue entry: %w", err) + } + return &email, nil +} + +func (r *emailQueueRepository) GetPendingEmails(ctx context.Context, limit int) ([]*models.EmailQueue, error) { + var emails []*models.EmailQueue + if err := r.db.WithContext(ctx). + Where("status = 'pending' AND scheduled_at <= ?", time.Now()). + Order("priority ASC, scheduled_at ASC"). + Limit(limit). + Find(&emails).Error; err != nil { + return nil, fmt.Errorf("failed to get pending emails: %w", err) + } + return emails, nil +} + +func (r *emailQueueRepository) Update(ctx context.Context, email *models.EmailQueue) error { + if err := r.db.WithContext(ctx).Save(email).Error; err != nil { + return fmt.Errorf("failed to update email queue entry: %w", err) + } + return nil +} + +func (r *emailQueueRepository) Delete(ctx context.Context, tenantID, emailID uuid.UUID) error { + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND id = ?", tenantID, emailID). + Delete(&models.EmailQueue{}).Error; err != nil { + return fmt.Errorf("failed to delete email queue entry: %w", err) + } + return nil +} + +func (r *emailQueueRepository) GetQueueStatus(ctx context.Context, tenantID uuid.UUID) (map[string]int, error) { + type statusCount struct { + Status string + Count int + } + var results []statusCount + + if err := r.db.WithContext(ctx). + Model(&models.EmailQueue{}). + Select("status, COUNT(*) as count"). + Where("tenant_id = ?", tenantID). + Group("status"). + Scan(&results).Error; err != nil { + return nil, fmt.Errorf("failed to get email queue status: %w", err) + } + + status := make(map[string]int) + for _, r := range results { + status[r.Status] = r.Count + } + return status, nil +} + +func (r *emailQueueRepository) GetFailedEmails(ctx context.Context, tenantID uuid.UUID, maxAttempts int) ([]*models.EmailQueue, error) { + var emails []*models.EmailQueue + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND status = 'failed' AND attempts < ?", tenantID, maxAttempts). + Order("created_at ASC"). + Find(&emails).Error; err != nil { + return nil, fmt.Errorf("failed to get failed emails: %w", err) + } + return emails, nil +} + +// ─── EmailVerification Repository ──────────────────────────────────────────── + +type emailVerificationRepository struct { + db *gorm.DB +} + +// NewEmailVerificationRepository creates a new email verification repository +func NewEmailVerificationRepository(db *gorm.DB) repo.EmailVerificationRepository { + return &emailVerificationRepository{db: db} +} + +func (r *emailVerificationRepository) Create(ctx context.Context, verification *models.EmailVerification) error { + if err := r.db.WithContext(ctx).Create(verification).Error; err != nil { + return fmt.Errorf("failed to create email verification: %w", err) + } + return nil +} + +func (r *emailVerificationRepository) GetByCode(ctx context.Context, tenantID uuid.UUID, code string) (*models.EmailVerification, error) { + var verification models.EmailVerification + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND code = ?", tenantID, code). + First(&verification).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrInvalidVerificationCode + } + return nil, fmt.Errorf("failed to get email verification by code: %w", err) + } + return &verification, nil +} + +func (r *emailVerificationRepository) GetByUserID(ctx context.Context, tenantID, userID uuid.UUID) (*models.EmailVerification, error) { + var verification models.EmailVerification + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND user_id = ? AND verified_at IS NULL", tenantID, userID). + Order("created_at DESC"). + First(&verification).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrResourceNotFound + } + return nil, fmt.Errorf("failed to get email verification by user ID: %w", err) + } + return &verification, nil +} + +func (r *emailVerificationRepository) Update(ctx context.Context, verification *models.EmailVerification) error { + if err := r.db.WithContext(ctx).Save(verification).Error; err != nil { + return fmt.Errorf("failed to update email verification: %w", err) + } + return nil +} + +func (r *emailVerificationRepository) Delete(ctx context.Context, tenantID uuid.UUID, code string) error { + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND code = ?", tenantID, code). + Delete(&models.EmailVerification{}).Error; err != nil { + return fmt.Errorf("failed to delete email verification: %w", err) + } + return nil +} + +func (r *emailVerificationRepository) DeleteExpired(ctx context.Context) error { + if err := r.db.WithContext(ctx). + Where("expires_at < ? AND verified_at IS NULL", time.Now()). + Delete(&models.EmailVerification{}).Error; err != nil { + return fmt.Errorf("failed to delete expired email verifications: %w", err) + } + return nil +} + +// ─── PasswordReset Repository ──────────────────────────────────────────────── + +type passwordResetRepository struct { + db *gorm.DB +} + +// NewPasswordResetRepository creates a new password reset repository +func NewPasswordResetRepository(db *gorm.DB) repo.PasswordResetRepository { + return &passwordResetRepository{db: db} +} + +func (r *passwordResetRepository) Create(ctx context.Context, reset *models.PasswordReset) error { + if err := r.db.WithContext(ctx).Create(reset).Error; err != nil { + return fmt.Errorf("failed to create password reset: %w", err) + } + return nil +} + +func (r *passwordResetRepository) GetByToken(ctx context.Context, tenantID uuid.UUID, token string) (*models.PasswordReset, error) { + var reset models.PasswordReset + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND token = ?", tenantID, token). + First(&reset).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrResourceNotFound + } + return nil, fmt.Errorf("failed to get password reset by token: %w", err) + } + return &reset, nil +} + +func (r *passwordResetRepository) GetByUserID(ctx context.Context, tenantID, userID uuid.UUID) ([]*models.PasswordReset, error) { + var resets []*models.PasswordReset + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND user_id = ?", tenantID, userID). + Order("created_at DESC"). + Find(&resets).Error; err != nil { + return nil, fmt.Errorf("failed to get password resets by user ID: %w", err) + } + return resets, nil +} + +func (r *passwordResetRepository) Update(ctx context.Context, reset *models.PasswordReset) error { + if err := r.db.WithContext(ctx).Save(reset).Error; err != nil { + return fmt.Errorf("failed to update password reset: %w", err) + } + return nil +} + +func (r *passwordResetRepository) Delete(ctx context.Context, tenantID uuid.UUID, token string) error { + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND token = ?", tenantID, token). + Delete(&models.PasswordReset{}).Error; err != nil { + return fmt.Errorf("failed to delete password reset: %w", err) + } + return nil +} + +func (r *passwordResetRepository) DeleteExpired(ctx context.Context) error { + if err := r.db.WithContext(ctx). + Where("expires_at < ? AND used_at IS NULL", time.Now()). + Delete(&models.PasswordReset{}).Error; err != nil { + return fmt.Errorf("failed to delete expired password resets: %w", err) + } + return nil +} diff --git a/internal/repo/gorm/permission.go b/internal/repo/gorm/permission.go new file mode 100644 index 0000000..8a058b3 --- /dev/null +++ b/internal/repo/gorm/permission.go @@ -0,0 +1,102 @@ +package gorm + +import ( + "context" + "errors" + "fmt" + + "shieldgate/internal/models" + "shieldgate/internal/repo" + + "github.com/google/uuid" + "gorm.io/gorm" +) + +type permissionRepository struct { + db *gorm.DB +} + +// NewPermissionRepository creates a new permission repository +func NewPermissionRepository(db *gorm.DB) repo.PermissionRepository { + return &permissionRepository{db: db} +} + +func (r *permissionRepository) Create(ctx context.Context, permission *models.Permission) error { + if err := r.db.WithContext(ctx).Create(permission).Error; err != nil { + return fmt.Errorf("failed to create permission: %w", err) + } + return nil +} + +func (r *permissionRepository) GetByID(ctx context.Context, permissionID uuid.UUID) (*models.Permission, error) { + var permission models.Permission + if err := r.db.WithContext(ctx). + Where("id = ?", permissionID). + First(&permission).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrPermissionNotFound + } + return nil, fmt.Errorf("failed to get permission by ID: %w", err) + } + return &permission, nil +} + +func (r *permissionRepository) GetByName(ctx context.Context, name string) (*models.Permission, error) { + var permission models.Permission + if err := r.db.WithContext(ctx). + Where("name = ?", name). + First(&permission).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrPermissionNotFound + } + return nil, fmt.Errorf("failed to get permission by name: %w", err) + } + return &permission, nil +} + +func (r *permissionRepository) Update(ctx context.Context, permission *models.Permission) error { + if err := r.db.WithContext(ctx).Save(permission).Error; err != nil { + return fmt.Errorf("failed to update permission: %w", err) + } + return nil +} + +func (r *permissionRepository) Delete(ctx context.Context, permissionID uuid.UUID) error { + if err := r.db.WithContext(ctx). + Delete(&models.Permission{}, "id = ?", permissionID).Error; err != nil { + return fmt.Errorf("failed to delete permission: %w", err) + } + return nil +} + +func (r *permissionRepository) List(ctx context.Context, limit, offset int) ([]*models.Permission, int64, error) { + var permissions []*models.Permission + var total int64 + + if err := r.db.WithContext(ctx). + Model(&models.Permission{}). + Count(&total).Error; err != nil { + return nil, 0, fmt.Errorf("failed to count permissions: %w", err) + } + + if err := r.db.WithContext(ctx). + Order("resource ASC, action ASC"). + Limit(limit). + Offset(offset). + Find(&permissions).Error; err != nil { + return nil, 0, fmt.Errorf("failed to list permissions: %w", err) + } + + return permissions, total, nil +} + +func (r *permissionRepository) GetByResource(ctx context.Context, resource string) ([]*models.Permission, error) { + var permissions []*models.Permission + if err := r.db.WithContext(ctx). + Where("resource = ?", resource). + Order("action ASC"). + Find(&permissions).Error; err != nil { + return nil, fmt.Errorf("failed to get permissions by resource: %w", err) + } + return permissions, nil +} diff --git a/internal/repo/gorm/repositories.go b/internal/repo/gorm/repositories.go index b5ef2e1..d93b95e 100644 --- a/internal/repo/gorm/repositories.go +++ b/internal/repo/gorm/repositories.go @@ -9,11 +9,27 @@ import ( // NewRepositories creates a new repositories instance with GORM implementations func NewRepositories(db *gorm.DB) *repo.Repositories { return &repo.Repositories{ + // Core OAuth Tenant: NewTenantRepository(db), User: NewUserRepository(db), Client: NewClientRepository(db), AuthCode: NewAuthCodeRepository(db), AccessToken: NewAccessTokenRepository(db), RefreshToken: NewRefreshTokenRepository(db), + + // RBAC + Role: NewRoleRepository(db), + Permission: NewPermissionRepository(db), + UserRole: NewUserRoleRepository(db), + RolePermission: NewRolePermissionRepository(db), + + // Audit + AuditLog: NewAuditLogRepository(db), + + // Email + EmailTemplate: NewEmailTemplateRepository(db), + EmailQueue: NewEmailQueueRepository(db), + EmailVerification: NewEmailVerificationRepository(db), + PasswordReset: NewPasswordResetRepository(db), } } diff --git a/internal/repo/gorm/role.go b/internal/repo/gorm/role.go new file mode 100644 index 0000000..e23052e --- /dev/null +++ b/internal/repo/gorm/role.go @@ -0,0 +1,108 @@ +package gorm + +import ( + "context" + "errors" + "fmt" + + "shieldgate/internal/models" + "shieldgate/internal/repo" + + "github.com/google/uuid" + "gorm.io/gorm" +) + +type roleRepository struct { + db *gorm.DB +} + +// NewRoleRepository creates a new role repository +func NewRoleRepository(db *gorm.DB) repo.RoleRepository { + return &roleRepository{db: db} +} + +func (r *roleRepository) Create(ctx context.Context, role *models.Role) error { + if err := r.db.WithContext(ctx).Create(role).Error; err != nil { + return fmt.Errorf("failed to create role: %w", err) + } + return nil +} + +func (r *roleRepository) GetByID(ctx context.Context, tenantID, roleID uuid.UUID) (*models.Role, error) { + var role models.Role + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND id = ?", tenantID, roleID). + First(&role).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrRoleNotFound + } + return nil, fmt.Errorf("failed to get role by ID: %w", err) + } + return &role, nil +} + +func (r *roleRepository) GetByName(ctx context.Context, tenantID uuid.UUID, name string) (*models.Role, error) { + var role models.Role + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND name = ?", tenantID, name). + First(&role).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrRoleNotFound + } + return nil, fmt.Errorf("failed to get role by name: %w", err) + } + return &role, nil +} + +func (r *roleRepository) Update(ctx context.Context, role *models.Role) error { + if err := r.db.WithContext(ctx).Save(role).Error; err != nil { + return fmt.Errorf("failed to update role: %w", err) + } + return nil +} + +func (r *roleRepository) Delete(ctx context.Context, tenantID, roleID uuid.UUID) error { + if err := r.db.WithContext(ctx). + Delete(&models.Role{}, "tenant_id = ? AND id = ?", tenantID, roleID).Error; err != nil { + return fmt.Errorf("failed to delete role: %w", err) + } + return nil +} + +func (r *roleRepository) List(ctx context.Context, tenantID uuid.UUID, limit, offset int) ([]*models.Role, int64, error) { + var roles []*models.Role + var total int64 + + if err := r.db.WithContext(ctx). + Model(&models.Role{}). + Where("tenant_id = ?", tenantID). + Count(&total).Error; err != nil { + return nil, 0, fmt.Errorf("failed to count roles: %w", err) + } + + if err := r.db.WithContext(ctx). + Where("tenant_id = ?", tenantID). + Order("created_at DESC"). + Limit(limit). + Offset(offset). + Find(&roles).Error; err != nil { + return nil, 0, fmt.Errorf("failed to list roles: %w", err) + } + + return roles, total, nil +} + +func (r *roleRepository) GetWithPermissions(ctx context.Context, tenantID, roleID uuid.UUID) (*models.Role, error) { + var role models.Role + if err := r.db.WithContext(ctx). + Preload("Permissions"). + Preload("Permissions.Permission"). + Where("tenant_id = ? AND id = ?", tenantID, roleID). + First(&role).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrRoleNotFound + } + return nil, fmt.Errorf("failed to get role with permissions: %w", err) + } + return &role, nil +} diff --git a/internal/repo/gorm/role_permission.go b/internal/repo/gorm/role_permission.go new file mode 100644 index 0000000..3a807a1 --- /dev/null +++ b/internal/repo/gorm/role_permission.go @@ -0,0 +1,86 @@ +package gorm + +import ( + "context" + "errors" + "fmt" + + "shieldgate/internal/models" + "shieldgate/internal/repo" + + "github.com/google/uuid" + "gorm.io/gorm" +) + +type rolePermissionRepository struct { + db *gorm.DB +} + +// NewRolePermissionRepository creates a new role-permission repository +func NewRolePermissionRepository(db *gorm.DB) repo.RolePermissionRepository { + return &rolePermissionRepository{db: db} +} + +func (r *rolePermissionRepository) Create(ctx context.Context, rolePermission *models.RolePermission) error { + if err := r.db.WithContext(ctx).Create(rolePermission).Error; err != nil { + return fmt.Errorf("failed to create role permission: %w", err) + } + return nil +} + +func (r *rolePermissionRepository) GetByID(ctx context.Context, rolePermissionID uuid.UUID) (*models.RolePermission, error) { + var rp models.RolePermission + if err := r.db.WithContext(ctx). + Where("id = ?", rolePermissionID). + First(&rp).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrResourceNotFound + } + return nil, fmt.Errorf("failed to get role permission by ID: %w", err) + } + return &rp, nil +} + +func (r *rolePermissionRepository) GetByRoleAndPermission(ctx context.Context, roleID, permissionID uuid.UUID) (*models.RolePermission, error) { + var rp models.RolePermission + if err := r.db.WithContext(ctx). + Where("role_id = ? AND permission_id = ?", roleID, permissionID). + First(&rp).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrResourceNotFound + } + return nil, fmt.Errorf("failed to get role permission: %w", err) + } + return &rp, nil +} + +func (r *rolePermissionRepository) Delete(ctx context.Context, roleID, permissionID uuid.UUID) error { + if err := r.db.WithContext(ctx). + Where("role_id = ? AND permission_id = ?", roleID, permissionID). + Delete(&models.RolePermission{}).Error; err != nil { + return fmt.Errorf("failed to delete role permission: %w", err) + } + return nil +} + +func (r *rolePermissionRepository) GetRolePermissions(ctx context.Context, roleID uuid.UUID) ([]*models.RolePermission, error) { + var rps []*models.RolePermission + if err := r.db.WithContext(ctx). + Preload("Permission"). + Where("role_id = ?", roleID). + Find(&rps).Error; err != nil { + return nil, fmt.Errorf("failed to get role permissions: %w", err) + } + return rps, nil +} + +func (r *rolePermissionRepository) GetPermissionRoles(ctx context.Context, permissionID uuid.UUID) ([]*models.RolePermission, error) { + var rps []*models.RolePermission + if err := r.db.WithContext(ctx). + Preload("Role"). + Where("permission_id = ?", permissionID). + Find(&rps).Error; err != nil { + return nil, fmt.Errorf("failed to get permission roles: %w", err) + } + return rps, nil +} diff --git a/internal/repo/gorm/user_role.go b/internal/repo/gorm/user_role.go new file mode 100644 index 0000000..1493ca8 --- /dev/null +++ b/internal/repo/gorm/user_role.go @@ -0,0 +1,98 @@ +package gorm + +import ( + "context" + "errors" + "fmt" + "time" + + "shieldgate/internal/models" + "shieldgate/internal/repo" + + "github.com/google/uuid" + "gorm.io/gorm" +) + +type userRoleRepository struct { + db *gorm.DB +} + +// NewUserRoleRepository creates a new user-role repository +func NewUserRoleRepository(db *gorm.DB) repo.UserRoleRepository { + return &userRoleRepository{db: db} +} + +func (r *userRoleRepository) Create(ctx context.Context, userRole *models.UserRole) error { + if err := r.db.WithContext(ctx).Create(userRole).Error; err != nil { + return fmt.Errorf("failed to create user role: %w", err) + } + return nil +} + +func (r *userRoleRepository) GetByID(ctx context.Context, tenantID, userRoleID uuid.UUID) (*models.UserRole, error) { + var userRole models.UserRole + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND id = ?", tenantID, userRoleID). + First(&userRole).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrResourceNotFound + } + return nil, fmt.Errorf("failed to get user role by ID: %w", err) + } + return &userRole, nil +} + +func (r *userRoleRepository) GetByUserAndRole(ctx context.Context, tenantID, userID, roleID uuid.UUID) (*models.UserRole, error) { + var userRole models.UserRole + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND user_id = ? AND role_id = ?", tenantID, userID, roleID). + First(&userRole).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, models.ErrResourceNotFound + } + return nil, fmt.Errorf("failed to get user role: %w", err) + } + return &userRole, nil +} + +func (r *userRoleRepository) Delete(ctx context.Context, tenantID, userID, roleID uuid.UUID) error { + if err := r.db.WithContext(ctx). + Where("tenant_id = ? AND user_id = ? AND role_id = ?", tenantID, userID, roleID). + Delete(&models.UserRole{}).Error; err != nil { + return fmt.Errorf("failed to delete user role: %w", err) + } + return nil +} + +func (r *userRoleRepository) GetUserRoles(ctx context.Context, tenantID, userID uuid.UUID) ([]*models.UserRole, error) { + var userRoles []*models.UserRole + if err := r.db.WithContext(ctx). + Preload("Role"). + Preload("Role.Permissions"). + Preload("Role.Permissions.Permission"). + Where("tenant_id = ? AND user_id = ?", tenantID, userID). + Find(&userRoles).Error; err != nil { + return nil, fmt.Errorf("failed to get user roles: %w", err) + } + return userRoles, nil +} + +func (r *userRoleRepository) GetRoleUsers(ctx context.Context, tenantID, roleID uuid.UUID) ([]*models.UserRole, error) { + var userRoles []*models.UserRole + if err := r.db.WithContext(ctx). + Preload("User"). + Where("tenant_id = ? AND role_id = ?", tenantID, roleID). + Find(&userRoles).Error; err != nil { + return nil, fmt.Errorf("failed to get role users: %w", err) + } + return userRoles, nil +} + +func (r *userRoleRepository) DeleteExpired(ctx context.Context) error { + if err := r.db.WithContext(ctx). + Where("expires_at IS NOT NULL AND expires_at < ?", time.Now()). + Delete(&models.UserRole{}).Error; err != nil { + return fmt.Errorf("failed to delete expired user roles: %w", err) + } + return nil +} diff --git a/internal/services/permission_service_impl.go b/internal/services/permission_service_impl.go new file mode 100644 index 0000000..e04e6aa --- /dev/null +++ b/internal/services/permission_service_impl.go @@ -0,0 +1,204 @@ +package services + +import ( + "context" + "fmt" + "time" + + "shieldgate/internal/models" + "shieldgate/internal/repo" + + "github.com/google/uuid" + "github.com/sirupsen/logrus" +) + +// PermissionServiceImpl implements the PermissionService interface +type PermissionServiceImpl struct { + permissionRepo repo.PermissionRepository + userRoleRepo repo.UserRoleRepository + rolePermissionRepo repo.RolePermissionRepository + logger *logrus.Logger +} + +// NewPermissionService creates a new permission service instance +func NewPermissionService( + permissionRepo repo.PermissionRepository, + userRoleRepo repo.UserRoleRepository, + rolePermissionRepo repo.RolePermissionRepository, + logger *logrus.Logger, +) PermissionService { + return &PermissionServiceImpl{ + permissionRepo: permissionRepo, + userRoleRepo: userRoleRepo, + rolePermissionRepo: rolePermissionRepo, + logger: logger, + } +} + +func (s *PermissionServiceImpl) Create(ctx context.Context, req *models.CreatePermissionRequest) (*models.Permission, error) { + s.logger.WithFields(logrus.Fields{ + "name": req.Name, + "resource": req.Resource, + "action": req.Action, + }).Info("creating new permission") + + existing, err := s.permissionRepo.GetByName(ctx, req.Name) + if err != nil && err != models.ErrPermissionNotFound { + return nil, fmt.Errorf("failed to check existing permission: %w", err) + } + if existing != nil { + return nil, models.ErrDuplicateResource + } + + permission := &models.Permission{ + ID: uuid.New(), + Name: req.Name, + DisplayName: req.DisplayName, + Description: req.Description, + Resource: req.Resource, + Action: req.Action, + IsSystem: false, + } + + if err := s.permissionRepo.Create(ctx, permission); err != nil { + return nil, fmt.Errorf("failed to create permission: %w", err) + } + + return permission, nil +} + +func (s *PermissionServiceImpl) GetByID(ctx context.Context, permissionID uuid.UUID) (*models.Permission, error) { + permission, err := s.permissionRepo.GetByID(ctx, permissionID) + if err != nil { + return nil, err + } + return permission, nil +} + +func (s *PermissionServiceImpl) GetByName(ctx context.Context, name string) (*models.Permission, error) { + permission, err := s.permissionRepo.GetByName(ctx, name) + if err != nil { + return nil, err + } + return permission, nil +} + +func (s *PermissionServiceImpl) Update(ctx context.Context, permissionID uuid.UUID, req *models.UpdatePermissionRequest) (*models.Permission, error) { + permission, err := s.permissionRepo.GetByID(ctx, permissionID) + if err != nil { + return nil, err + } + + if permission.IsSystem { + return nil, fmt.Errorf("system permissions cannot be modified: %w", models.ErrBusinessRuleViolation) + } + + if req.DisplayName != "" { + permission.DisplayName = req.DisplayName + } + if req.Description != "" { + permission.Description = req.Description + } + if req.Resource != "" { + permission.Resource = req.Resource + } + if req.Action != "" { + permission.Action = req.Action + } + + if err := s.permissionRepo.Update(ctx, permission); err != nil { + return nil, fmt.Errorf("failed to update permission: %w", err) + } + + return permission, nil +} + +func (s *PermissionServiceImpl) Delete(ctx context.Context, permissionID uuid.UUID) error { + permission, err := s.permissionRepo.GetByID(ctx, permissionID) + if err != nil { + return err + } + + if permission.IsSystem { + return fmt.Errorf("system permissions cannot be deleted: %w", models.ErrBusinessRuleViolation) + } + + return s.permissionRepo.Delete(ctx, permissionID) +} + +func (s *PermissionServiceImpl) List(ctx context.Context, limit, offset int) (*models.PaginatedResponse, error) { + permissions, total, err := s.permissionRepo.List(ctx, limit, offset) + if err != nil { + return nil, fmt.Errorf("failed to list permissions: %w", err) + } + + items := make([]interface{}, len(permissions)) + for i, p := range permissions { + items[i] = p + } + + return models.NewPaginatedResponse(items, limit, offset, total), nil +} + +// HasPermission checks whether a user has the given resource+action permission +// by traversing their assigned roles → role permissions chain. +func (s *PermissionServiceImpl) HasPermission(ctx context.Context, tenantID, userID uuid.UUID, resource, action string) (bool, error) { + userRoles, err := s.userRoleRepo.GetUserRoles(ctx, tenantID, userID) + if err != nil { + return false, fmt.Errorf("failed to get user roles: %w", err) + } + + for _, ur := range userRoles { + // Skip expired role assignments + if ur.ExpiresAt != nil && time.Now().After(*ur.ExpiresAt) { + continue + } + + rolePerms, err := s.rolePermissionRepo.GetRolePermissions(ctx, ur.RoleID) + if err != nil { + s.logger.WithError(err).WithField("role_id", ur.RoleID).Warn("failed to get role permissions") + continue + } + + for _, rp := range rolePerms { + if rp.Permission.Resource == resource && rp.Permission.Action == action { + return true, nil + } + // "manage" action grants all actions on the resource + if rp.Permission.Resource == resource && rp.Permission.Action == "manage" { + return true, nil + } + } + } + + return false, nil +} + +// GetUserPermissions returns all permissions a user has across all their roles +func (s *PermissionServiceImpl) GetUserPermissions(ctx context.Context, tenantID, userID uuid.UUID) ([]*models.Permission, error) { + userRoles, err := s.userRoleRepo.GetUserRoles(ctx, tenantID, userID) + if err != nil { + return nil, fmt.Errorf("failed to get user roles: %w", err) + } + + seen := make(map[uuid.UUID]struct{}) + var permissions []*models.Permission + + for _, ur := range userRoles { + rolePerms, err := s.rolePermissionRepo.GetRolePermissions(ctx, ur.RoleID) + if err != nil { + s.logger.WithError(err).WithField("role_id", ur.RoleID).Warn("failed to get role permissions") + continue + } + + for _, rp := range rolePerms { + if _, exists := seen[rp.PermissionID]; !exists { + seen[rp.PermissionID] = struct{}{} + p := rp.Permission + permissions = append(permissions, &p) + } + } + } + + return permissions, nil +} diff --git a/internal/services/tests/permission_service_test.go b/internal/services/tests/permission_service_test.go new file mode 100644 index 0000000..3474d03 --- /dev/null +++ b/internal/services/tests/permission_service_test.go @@ -0,0 +1,393 @@ +package tests + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/google/uuid" + "github.com/sirupsen/logrus" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "shieldgate/internal/models" + "shieldgate/internal/services" +) + +// ─── Mock Repositories ─────────────────────────────────────────────────────── + +type mockPermissionRepo struct { + permissions map[uuid.UUID]*models.Permission + byName map[string]*models.Permission +} + +func newMockPermissionRepo() *mockPermissionRepo { + return &mockPermissionRepo{ + permissions: make(map[uuid.UUID]*models.Permission), + byName: make(map[string]*models.Permission), + } +} + +func (m *mockPermissionRepo) Create(_ context.Context, p *models.Permission) error { + m.permissions[p.ID] = p + m.byName[p.Name] = p + return nil +} + +func (m *mockPermissionRepo) GetByID(_ context.Context, id uuid.UUID) (*models.Permission, error) { + p, ok := m.permissions[id] + if !ok { + return nil, models.ErrPermissionNotFound + } + return p, nil +} + +func (m *mockPermissionRepo) GetByName(_ context.Context, name string) (*models.Permission, error) { + p, ok := m.byName[name] + if !ok { + return nil, models.ErrPermissionNotFound + } + return p, nil +} + +func (m *mockPermissionRepo) Update(_ context.Context, p *models.Permission) error { + m.permissions[p.ID] = p + m.byName[p.Name] = p + return nil +} + +func (m *mockPermissionRepo) Delete(_ context.Context, id uuid.UUID) error { + if p, ok := m.permissions[id]; ok { + delete(m.byName, p.Name) + delete(m.permissions, id) + } + return nil +} + +func (m *mockPermissionRepo) List(_ context.Context, limit, offset int) ([]*models.Permission, int64, error) { + var all []*models.Permission + for _, p := range m.permissions { + all = append(all, p) + } + total := int64(len(all)) + end := offset + limit + if end > len(all) { + end = len(all) + } + if offset >= len(all) { + return nil, total, nil + } + return all[offset:end], total, nil +} + +func (m *mockPermissionRepo) GetByResource(_ context.Context, resource string) ([]*models.Permission, error) { + var result []*models.Permission + for _, p := range m.permissions { + if p.Resource == resource { + result = append(result, p) + } + } + return result, nil +} + +type mockUserRoleRepo struct { + userRoles []*models.UserRole +} + +func (m *mockUserRoleRepo) Create(_ context.Context, ur *models.UserRole) error { + m.userRoles = append(m.userRoles, ur) + return nil +} + +func (m *mockUserRoleRepo) GetByID(_ context.Context, _, id uuid.UUID) (*models.UserRole, error) { + for _, ur := range m.userRoles { + if ur.ID == id { + return ur, nil + } + } + return nil, models.ErrResourceNotFound +} + +func (m *mockUserRoleRepo) GetByUserAndRole(_ context.Context, tenantID, userID, roleID uuid.UUID) (*models.UserRole, error) { + for _, ur := range m.userRoles { + if ur.TenantID == tenantID && ur.UserID == userID && ur.RoleID == roleID { + return ur, nil + } + } + return nil, models.ErrResourceNotFound +} + +func (m *mockUserRoleRepo) Delete(_ context.Context, tenantID, userID, roleID uuid.UUID) error { + for i, ur := range m.userRoles { + if ur.TenantID == tenantID && ur.UserID == userID && ur.RoleID == roleID { + m.userRoles = append(m.userRoles[:i], m.userRoles[i+1:]...) + return nil + } + } + return nil +} + +func (m *mockUserRoleRepo) GetUserRoles(_ context.Context, tenantID, userID uuid.UUID) ([]*models.UserRole, error) { + var result []*models.UserRole + for _, ur := range m.userRoles { + if ur.TenantID == tenantID && ur.UserID == userID { + result = append(result, ur) + } + } + return result, nil +} + +func (m *mockUserRoleRepo) GetRoleUsers(_ context.Context, tenantID, roleID uuid.UUID) ([]*models.UserRole, error) { + var result []*models.UserRole + for _, ur := range m.userRoles { + if ur.TenantID == tenantID && ur.RoleID == roleID { + result = append(result, ur) + } + } + return result, nil +} + +func (m *mockUserRoleRepo) DeleteExpired(_ context.Context) error { return nil } + +type mockRolePermRepo struct { + rolePerms []*models.RolePermission +} + +func (m *mockRolePermRepo) Create(_ context.Context, rp *models.RolePermission) error { + m.rolePerms = append(m.rolePerms, rp) + return nil +} + +func (m *mockRolePermRepo) GetByID(_ context.Context, id uuid.UUID) (*models.RolePermission, error) { + for _, rp := range m.rolePerms { + if rp.ID == id { + return rp, nil + } + } + return nil, models.ErrResourceNotFound +} + +func (m *mockRolePermRepo) GetByRoleAndPermission(_ context.Context, roleID, permID uuid.UUID) (*models.RolePermission, error) { + for _, rp := range m.rolePerms { + if rp.RoleID == roleID && rp.PermissionID == permID { + return rp, nil + } + } + return nil, models.ErrResourceNotFound +} + +func (m *mockRolePermRepo) Delete(_ context.Context, roleID, permID uuid.UUID) error { + for i, rp := range m.rolePerms { + if rp.RoleID == roleID && rp.PermissionID == permID { + m.rolePerms = append(m.rolePerms[:i], m.rolePerms[i+1:]...) + return nil + } + } + return nil +} + +func (m *mockRolePermRepo) GetRolePermissions(_ context.Context, roleID uuid.UUID) ([]*models.RolePermission, error) { + var result []*models.RolePermission + for _, rp := range m.rolePerms { + if rp.RoleID == roleID { + result = append(result, rp) + } + } + return result, nil +} + +func (m *mockRolePermRepo) GetPermissionRoles(_ context.Context, permID uuid.UUID) ([]*models.RolePermission, error) { + var result []*models.RolePermission + for _, rp := range m.rolePerms { + if rp.PermissionID == permID { + result = append(result, rp) + } + } + return result, nil +} + +// ─── Helper ────────────────────────────────────────────────────────────────── + +func newTestPermissionService(permRepo *mockPermissionRepo, urRepo *mockUserRoleRepo, rpRepo *mockRolePermRepo) services.PermissionService { + logger := logrus.New() + logger.SetLevel(logrus.FatalLevel) + return services.NewPermissionService(permRepo, urRepo, rpRepo, logger) +} + +// ─── Tests ──────────────────────────────────────────────────────────────────── + +func TestPermissionService_Create_Success(t *testing.T) { + svc := newTestPermissionService(newMockPermissionRepo(), &mockUserRoleRepo{}, &mockRolePermRepo{}) + + req := &models.CreatePermissionRequest{ + Name: "users.read", + DisplayName: "Read Users", + Resource: "user", + Action: "read", + } + p, err := svc.Create(context.Background(), req) + + require.NoError(t, err) + assert.Equal(t, "users.read", p.Name) + assert.Equal(t, "user", p.Resource) + assert.Equal(t, "read", p.Action) +} + +func TestPermissionService_Create_Duplicate(t *testing.T) { + permRepo := newMockPermissionRepo() + svc := newTestPermissionService(permRepo, &mockUserRoleRepo{}, &mockRolePermRepo{}) + + req := &models.CreatePermissionRequest{ + Name: "users.read", + DisplayName: "Read Users", + Resource: "user", + Action: "read", + } + _, err := svc.Create(context.Background(), req) + require.NoError(t, err) + + _, err = svc.Create(context.Background(), req) + assert.True(t, errors.Is(err, models.ErrDuplicateResource)) +} + +func TestPermissionService_Delete_SystemPermission_Rejected(t *testing.T) { + permRepo := newMockPermissionRepo() + id := uuid.New() + systemPerm := &models.Permission{ + ID: id, + Name: "system.perm", + Resource: "system", + Action: "manage", + IsSystem: true, + } + permRepo.permissions[id] = systemPerm + permRepo.byName["system.perm"] = systemPerm + + svc := newTestPermissionService(permRepo, &mockUserRoleRepo{}, &mockRolePermRepo{}) + + err := svc.Delete(context.Background(), id) + assert.Error(t, err, "system permissions should not be deletable") +} + +func TestPermissionService_HasPermission_ExactMatch(t *testing.T) { + permRepo := newMockPermissionRepo() + urRepo := &mockUserRoleRepo{} + rpRepo := &mockRolePermRepo{} + + tenantID := uuid.New() + userID := uuid.New() + roleID := uuid.New() + permID := uuid.New() + + perm := &models.Permission{ID: permID, Name: "users.read", Resource: "user", Action: "read"} + permRepo.permissions[permID] = perm + + urRepo.userRoles = []*models.UserRole{ + {ID: uuid.New(), TenantID: tenantID, UserID: userID, RoleID: roleID, GrantedAt: time.Now()}, + } + rpRepo.rolePerms = []*models.RolePermission{ + {ID: uuid.New(), RoleID: roleID, PermissionID: permID, Permission: *perm}, + } + + svc := newTestPermissionService(permRepo, urRepo, rpRepo) + + has, err := svc.HasPermission(context.Background(), tenantID, userID, "user", "read") + require.NoError(t, err) + assert.True(t, has) +} + +func TestPermissionService_HasPermission_ManageGrantsAll(t *testing.T) { + permRepo := newMockPermissionRepo() + urRepo := &mockUserRoleRepo{} + rpRepo := &mockRolePermRepo{} + + tenantID := uuid.New() + userID := uuid.New() + roleID := uuid.New() + permID := uuid.New() + + managePerm := &models.Permission{ID: permID, Name: "users.manage", Resource: "user", Action: "manage"} + permRepo.permissions[permID] = managePerm + + urRepo.userRoles = []*models.UserRole{ + {ID: uuid.New(), TenantID: tenantID, UserID: userID, RoleID: roleID, GrantedAt: time.Now()}, + } + rpRepo.rolePerms = []*models.RolePermission{ + {ID: uuid.New(), RoleID: roleID, PermissionID: permID, Permission: *managePerm}, + } + + svc := newTestPermissionService(permRepo, urRepo, rpRepo) + + // "manage" should grant "delete" + has, err := svc.HasPermission(context.Background(), tenantID, userID, "user", "delete") + require.NoError(t, err) + assert.True(t, has) +} + +func TestPermissionService_HasPermission_ExpiredRole_Denied(t *testing.T) { + permRepo := newMockPermissionRepo() + urRepo := &mockUserRoleRepo{} + rpRepo := &mockRolePermRepo{} + + tenantID := uuid.New() + userID := uuid.New() + roleID := uuid.New() + permID := uuid.New() + + perm := &models.Permission{ID: permID, Name: "users.read", Resource: "user", Action: "read"} + permRepo.permissions[permID] = perm + + past := time.Now().Add(-1 * time.Hour) + urRepo.userRoles = []*models.UserRole{ + {ID: uuid.New(), TenantID: tenantID, UserID: userID, RoleID: roleID, ExpiresAt: &past}, + } + rpRepo.rolePerms = []*models.RolePermission{ + {ID: uuid.New(), RoleID: roleID, PermissionID: permID, Permission: *perm}, + } + + svc := newTestPermissionService(permRepo, urRepo, rpRepo) + + has, err := svc.HasPermission(context.Background(), tenantID, userID, "user", "read") + require.NoError(t, err) + assert.False(t, has, "expired role should not grant permissions") +} + +func TestPermissionService_HasPermission_NoRoles_Denied(t *testing.T) { + svc := newTestPermissionService(newMockPermissionRepo(), &mockUserRoleRepo{}, &mockRolePermRepo{}) + + has, err := svc.HasPermission(context.Background(), uuid.New(), uuid.New(), "user", "read") + require.NoError(t, err) + assert.False(t, has) +} + +func TestPermissionService_GetUserPermissions_DeduplicatesAcrossRoles(t *testing.T) { + permRepo := newMockPermissionRepo() + urRepo := &mockUserRoleRepo{} + rpRepo := &mockRolePermRepo{} + + tenantID := uuid.New() + userID := uuid.New() + roleID1 := uuid.New() + roleID2 := uuid.New() + permID := uuid.New() + + perm := &models.Permission{ID: permID, Name: "users.read", Resource: "user", Action: "read"} + permRepo.permissions[permID] = perm + + urRepo.userRoles = []*models.UserRole{ + {ID: uuid.New(), TenantID: tenantID, UserID: userID, RoleID: roleID1}, + {ID: uuid.New(), TenantID: tenantID, UserID: userID, RoleID: roleID2}, + } + // Both roles have the same permission — should appear once in result + rpRepo.rolePerms = []*models.RolePermission{ + {ID: uuid.New(), RoleID: roleID1, PermissionID: permID, Permission: *perm}, + {ID: uuid.New(), RoleID: roleID2, PermissionID: permID, Permission: *perm}, + } + + svc := newTestPermissionService(permRepo, urRepo, rpRepo) + + perms, err := svc.GetUserPermissions(context.Background(), tenantID, userID) + require.NoError(t, err) + assert.Len(t, perms, 1, "duplicate permissions from multiple roles should be deduplicated") +}