mengstack-api/internal/modules/notification/infrastructure/notification_repo.go
MengStack Dev cb7ad4d35d feat(notification): M6 示例模块 + 单元测试
Complete 4-layer notification module as reference template:
domain → infrastructure → application → interfaces.
10 unit tests covering all service methods with mock repository.
2026-10-03 01:07:10 +08:00

72 lines
2.2 KiB
Go

package infrastructure
import (
"context"
"mengstack/internal/modules/notification/domain"
"gorm.io/gorm"
)
type notificationRepo struct {
db *gorm.DB
}
func NewNotificationRepository(db *gorm.DB) domain.NotificationRepository {
return &notificationRepo{db: db}
}
func (r *notificationRepo) Create(ctx context.Context, n *domain.Notification) error {
return r.db.WithContext(ctx).Create(n).Error
}
func (r *notificationRepo) FindByID(ctx context.Context, tenantID string, id uint) (*domain.Notification, error) {
var n domain.Notification
err := r.db.WithContext(ctx).Where("id = ? AND tenant_id = ?", id, tenantID).First(&n).Error
if err != nil {
return nil, err
}
return &n, nil
}
func (r *notificationRepo) FindByUser(ctx context.Context, tenantID string, userID uint, page, pageSize int) ([]domain.Notification, int64, error) {
var notifications []domain.Notification
var total int64
query := r.db.WithContext(ctx).Where("tenant_id = ? AND user_id = ?", tenantID, userID)
query.Model(&domain.Notification{}).Count(&total)
offset := (page - 1) * pageSize
err := query.Order("created_at DESC").Offset(offset).Limit(pageSize).Find(&notifications).Error
return notifications, total, err
}
func (r *notificationRepo) MarkAsRead(ctx context.Context, tenantID string, id uint) error {
return r.db.WithContext(ctx).
Where("id = ? AND tenant_id = ?", id, tenantID).
Model(&domain.Notification{}).
Update("is_read", true).Error
}
func (r *notificationRepo) MarkAllAsRead(ctx context.Context, tenantID string, userID uint) error {
return r.db.WithContext(ctx).
Where("tenant_id = ? AND user_id = ? AND is_read = ?", tenantID, userID, false).
Model(&domain.Notification{}).
Update("is_read", true).Error
}
func (r *notificationRepo) CountUnread(ctx context.Context, tenantID string, userID uint) (int64, error) {
var count int64
err := r.db.WithContext(ctx).
Where("tenant_id = ? AND user_id = ? AND is_read = ?", tenantID, userID, false).
Model(&domain.Notification{}).
Count(&count).Error
return count, err
}
func (r *notificationRepo) Delete(ctx context.Context, tenantID string, id uint) error {
return r.db.WithContext(ctx).
Where("id = ? AND tenant_id = ?", id, tenantID).
Delete(&domain.Notification{}).Error
}