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.
This commit is contained in:
MengStack Dev 2026-10-03 01:07:10 +08:00
parent 1f4e2ca0d0
commit cb7ad4d35d
9 changed files with 677 additions and 0 deletions

View File

@ -17,6 +17,7 @@ import (
orginterfaces "mengstack/internal/modules/org/interfaces"
auditinterfaces "mengstack/internal/modules/audit/interfaces"
settingsinterfaces "mengstack/internal/modules/settings/interfaces"
notificationinterfaces "mengstack/internal/modules/notification/interfaces"
"github.com/gin-gonic/gin"
"github.com/redis/go-redis/v9"
@ -45,6 +46,7 @@ func newEngine(
orgHandler *orginterfaces.Handler,
auditHandler *auditinterfaces.Handler,
settingsHandler *settingsinterfaces.Handler,
notificationHandler *notificationinterfaces.Handler,
) *gin.Engine {
ginMode := "release"
if cfg.Server.Mode == "debug" || cfg.Server.Mode == "dev" {
@ -68,6 +70,7 @@ func newEngine(
orginterfaces.SetupRoutes(r, orgHandler, authMW)
auditinterfaces.SetupRoutes(r, auditHandler, authMW)
settingsinterfaces.SetupRoutes(r, settingsHandler, authMW)
notificationinterfaces.SetupRoutes(r, notificationHandler, authMW)
return r
}
@ -113,6 +116,7 @@ func NewApp() *fx.App {
orginterfaces.Module,
auditinterfaces.Module,
settingsinterfaces.Module,
notificationinterfaces.Module,
Module,
)
}

View File

@ -0,0 +1,82 @@
package application
import (
"context"
"errors"
"mengstack/internal/modules/notification/domain"
"gorm.io/gorm"
)
type Service struct {
repo domain.NotificationRepository
}
func NewService(repo domain.NotificationRepository) *Service {
return &Service{repo: repo}
}
func (s *Service) Create(ctx context.Context, tenantID string, req *domain.CreateNotificationRequest) (*domain.NotificationDTO, error) {
n := &domain.Notification{
TenantID: tenantID,
UserID: req.UserID,
Title: req.Title,
Content: req.Content,
Type: req.Type,
}
if n.Type == "" {
n.Type = "info"
}
if err := s.repo.Create(ctx, n); err != nil {
return nil, err
}
dto := domain.ToNotificationDTO(n)
return &dto, nil
}
func (s *Service) Get(ctx context.Context, tenantID string, id uint) (*domain.NotificationDTO, error) {
n, err := s.repo.FindByID(ctx, tenantID, id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
return nil, err
}
dto := domain.ToNotificationDTO(n)
return &dto, nil
}
func (s *Service) ListByUser(ctx context.Context, tenantID string, userID uint, page, pageSize int) ([]domain.NotificationDTO, int64, error) {
if page < 1 {
page = 1
}
if pageSize < 1 || pageSize > 100 {
pageSize = 20
}
notifications, total, err := s.repo.FindByUser(ctx, tenantID, userID, page, pageSize)
if err != nil {
return nil, 0, err
}
result := make([]domain.NotificationDTO, len(notifications))
for i, n := range notifications {
result[i] = domain.ToNotificationDTO(&n)
}
return result, total, nil
}
func (s *Service) MarkAsRead(ctx context.Context, tenantID string, id uint) error {
return s.repo.MarkAsRead(ctx, tenantID, id)
}
func (s *Service) MarkAllAsRead(ctx context.Context, tenantID string, userID uint) error {
return s.repo.MarkAllAsRead(ctx, tenantID, userID)
}
func (s *Service) CountUnread(ctx context.Context, tenantID string, userID uint) (int64, error) {
return s.repo.CountUnread(ctx, tenantID, userID)
}
func (s *Service) Delete(ctx context.Context, tenantID string, id uint) error {
return s.repo.Delete(ctx, tenantID, id)
}

View File

@ -0,0 +1,279 @@
package application
import (
"context"
"testing"
"mengstack/internal/modules/notification/domain"
"gorm.io/gorm"
)
type mockRepo struct {
notifications []domain.Notification
nextID uint
}
func newMockRepo() *mockRepo {
return &mockRepo{nextID: 1}
}
func (m *mockRepo) Create(_ context.Context, n *domain.Notification) error {
n.ID = m.nextID
m.nextID++
m.notifications = append(m.notifications, *n)
return nil
}
func (m *mockRepo) FindByID(_ context.Context, tenantID string, id uint) (*domain.Notification, error) {
for _, n := range m.notifications {
if n.ID == id && n.TenantID == tenantID {
return &n, nil
}
}
return nil, gorm.ErrRecordNotFound
}
func (m *mockRepo) FindByUser(_ context.Context, tenantID string, userID uint, page, pageSize int) ([]domain.Notification, int64, error) {
var result []domain.Notification
for _, n := range m.notifications {
if n.TenantID == tenantID && n.UserID == userID {
result = append(result, n)
}
}
total := int64(len(result))
start := (page - 1) * pageSize
if start >= len(result) {
return nil, total, nil
}
end := start + pageSize
if end > len(result) {
end = len(result)
}
return result[start:end], total, nil
}
func (m *mockRepo) MarkAsRead(_ context.Context, tenantID string, id uint) error {
for i, n := range m.notifications {
if n.ID == id && n.TenantID == tenantID {
m.notifications[i].IsRead = true
return nil
}
}
return nil
}
func (m *mockRepo) MarkAllAsRead(_ context.Context, tenantID string, userID uint) error {
for i, n := range m.notifications {
if n.TenantID == tenantID && n.UserID == userID {
m.notifications[i].IsRead = true
}
}
return nil
}
func (m *mockRepo) CountUnread(_ context.Context, tenantID string, userID uint) (int64, error) {
var count int64
for _, n := range m.notifications {
if n.TenantID == tenantID && n.UserID == userID && !n.IsRead {
count++
}
}
return count, nil
}
func (m *mockRepo) Delete(_ context.Context, tenantID string, id uint) error {
for i, n := range m.notifications {
if n.ID == id && n.TenantID == tenantID {
m.notifications = append(m.notifications[:i], m.notifications[i+1:]...)
return nil
}
}
return nil
}
func TestCreate(t *testing.T) {
repo := newMockRepo()
svc := NewService(repo)
req := &domain.CreateNotificationRequest{
UserID: 1,
Title: "Test Notification",
Content: "Hello",
Type: "warning",
}
dto, err := svc.Create(context.Background(), "tenant-1", req)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if dto.ID != 1 {
t.Errorf("expected ID 1, got %d", dto.ID)
}
if dto.Title != "Test Notification" {
t.Errorf("expected title 'Test Notification', got %q", dto.Title)
}
if dto.Type != "warning" {
t.Errorf("expected type 'warning', got %q", dto.Type)
}
}
func TestCreateDefaultType(t *testing.T) {
repo := newMockRepo()
svc := NewService(repo)
req := &domain.CreateNotificationRequest{
UserID: 1,
Title: "Test",
}
dto, err := svc.Create(context.Background(), "tenant-1", req)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if dto.Type != "info" {
t.Errorf("expected default type 'info', got %q", dto.Type)
}
}
func TestGet(t *testing.T) {
repo := newMockRepo()
svc := NewService(repo)
req := &domain.CreateNotificationRequest{UserID: 1, Title: "Test"}
svc.Create(context.Background(), "tenant-1", req)
dto, err := svc.Get(context.Background(), "tenant-1", 1)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if dto == nil {
t.Fatal("expected notification, got nil")
}
if dto.Title != "Test" {
t.Errorf("expected title 'Test', got %q", dto.Title)
}
}
func TestGetNotFound(t *testing.T) {
repo := newMockRepo()
svc := NewService(repo)
dto, err := svc.Get(context.Background(), "tenant-1", 999)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if dto != nil {
t.Error("expected nil for non-existent notification")
}
}
func TestListByUser(t *testing.T) {
repo := newMockRepo()
svc := NewService(repo)
for i := 0; i < 5; i++ {
svc.Create(context.Background(), "tenant-1", &domain.CreateNotificationRequest{UserID: 1, Title: "Test"})
}
svc.Create(context.Background(), "tenant-1", &domain.CreateNotificationRequest{UserID: 2, Title: "Other user"})
items, total, err := svc.ListByUser(context.Background(), "tenant-1", 1, 1, 10)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if total != 5 {
t.Errorf("expected total 5, got %d", total)
}
if len(items) != 5 {
t.Errorf("expected 5 items, got %d", len(items))
}
}
func TestListByUserPagination(t *testing.T) {
repo := newMockRepo()
svc := NewService(repo)
for i := 0; i < 5; i++ {
svc.Create(context.Background(), "tenant-1", &domain.CreateNotificationRequest{UserID: 1, Title: "Test"})
}
items, total, err := svc.ListByUser(context.Background(), "tenant-1", 1, 1, 2)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if total != 5 {
t.Errorf("expected total 5, got %d", total)
}
if len(items) != 2 {
t.Errorf("expected 2 items on page 1, got %d", len(items))
}
}
func TestMarkAsRead(t *testing.T) {
repo := newMockRepo()
svc := NewService(repo)
svc.Create(context.Background(), "tenant-1", &domain.CreateNotificationRequest{UserID: 1, Title: "Test"})
err := svc.MarkAsRead(context.Background(), "tenant-1", 1)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
dto, _ := svc.Get(context.Background(), "tenant-1", 1)
if !dto.IsRead {
t.Error("expected notification to be marked as read")
}
}
func TestMarkAllAsRead(t *testing.T) {
repo := newMockRepo()
svc := NewService(repo)
for i := 0; i < 3; i++ {
svc.Create(context.Background(), "tenant-1", &domain.CreateNotificationRequest{UserID: 1, Title: "Test"})
}
err := svc.MarkAllAsRead(context.Background(), "tenant-1", 1)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
count, _ := svc.CountUnread(context.Background(), "tenant-1", 1)
if count != 0 {
t.Errorf("expected 0 unread, got %d", count)
}
}
func TestCountUnread(t *testing.T) {
repo := newMockRepo()
svc := NewService(repo)
for i := 0; i < 3; i++ {
svc.Create(context.Background(), "tenant-1", &domain.CreateNotificationRequest{UserID: 1, Title: "Test"})
}
svc.MarkAsRead(context.Background(), "tenant-1", 1)
count, err := svc.CountUnread(context.Background(), "tenant-1", 1)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if count != 2 {
t.Errorf("expected 2 unread, got %d", count)
}
}
func TestDelete(t *testing.T) {
repo := newMockRepo()
svc := NewService(repo)
svc.Create(context.Background(), "tenant-1", &domain.CreateNotificationRequest{UserID: 1, Title: "Test"})
err := svc.Delete(context.Background(), "tenant-1", 1)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
dto, _ := svc.Get(context.Background(), "tenant-1", 1)
if dto != nil {
t.Error("expected nil after delete")
}
}

View File

@ -0,0 +1,50 @@
package domain
import "time"
type Notification struct {
ID uint `json:"id" gorm:"primaryKey"`
TenantID string `json:"tenant_id" gorm:"size:36;not null;index"`
UserID uint `json:"user_id" gorm:"not null;index"`
Title string `json:"title" gorm:"size:256;not null"`
Content string `json:"content" gorm:"type:text"`
Type string `json:"type" gorm:"size:32;default:info"`
IsRead bool `json:"is_read" gorm:"default:false"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func (Notification) TableName() string {
return "notifications"
}
type NotificationDTO struct {
ID uint `json:"id"`
TenantID string `json:"tenant_id"`
UserID uint `json:"user_id"`
Title string `json:"title"`
Content string `json:"content"`
Type string `json:"type"`
IsRead bool `json:"is_read"`
CreatedAt time.Time `json:"created_at"`
}
func ToNotificationDTO(n *Notification) NotificationDTO {
return NotificationDTO{
ID: n.ID,
TenantID: n.TenantID,
UserID: n.UserID,
Title: n.Title,
Content: n.Content,
Type: n.Type,
IsRead: n.IsRead,
CreatedAt: n.CreatedAt,
}
}
type CreateNotificationRequest struct {
UserID uint `json:"user_id" binding:"required"`
Title string `json:"title" binding:"required"`
Content string `json:"content"`
Type string `json:"type"`
}

View File

@ -0,0 +1,13 @@
package domain
import "context"
type NotificationRepository interface {
Create(ctx context.Context, notification *Notification) error
FindByID(ctx context.Context, tenantID string, id uint) (*Notification, error)
FindByUser(ctx context.Context, tenantID string, userID uint, page, pageSize int) ([]Notification, int64, error)
MarkAsRead(ctx context.Context, tenantID string, id uint) error
MarkAllAsRead(ctx context.Context, tenantID string, userID uint) error
CountUnread(ctx context.Context, tenantID string, userID uint) (int64, error)
Delete(ctx context.Context, tenantID string, id uint) error
}

View File

@ -0,0 +1,11 @@
package infrastructure
import (
"mengstack/internal/modules/notification/domain"
"gorm.io/gorm"
)
func Migrate(db *gorm.DB) error {
return db.AutoMigrate(&domain.Notification{})
}

View File

@ -0,0 +1,71 @@
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
}

View File

@ -0,0 +1,124 @@
package interfaces
import (
"net/http"
"strconv"
"mengstack/internal/kernel/response"
"mengstack/internal/modules/notification/application"
"mengstack/internal/modules/notification/domain"
"github.com/gin-gonic/gin"
)
type Handler struct {
svc *application.Service
}
func NewHandler(svc *application.Service) *Handler {
return &Handler{svc: svc}
}
func (h *Handler) Create(c *gin.Context) {
var req domain.CreateNotificationRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, response.NewBadRequest("invalid request: "+err.Error()))
return
}
tenantID := c.GetString("tenant_id")
n, err := h.svc.Create(c.Request.Context(), tenantID, &req)
if err != nil {
response.HandleError(c, err)
return
}
response.Success(c, n)
}
func (h *Handler) Get(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, response.NewBadRequest("invalid id"))
return
}
tenantID := c.GetString("tenant_id")
n, err := h.svc.Get(c.Request.Context(), tenantID, uint(id))
if err != nil {
response.HandleError(c, err)
return
}
if n == nil {
response.Fail(c, response.NewBadRequest("notification not found"))
return
}
response.Success(c, n)
}
func (h *Handler) ListByUser(c *gin.Context) {
userID, err := strconv.ParseUint(c.Param("userId"), 10, 64)
if err != nil {
response.Fail(c, response.NewBadRequest("invalid user_id"))
return
}
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
tenantID := c.GetString("tenant_id")
notifications, total, err := h.svc.ListByUser(c.Request.Context(), tenantID, uint(userID), page, pageSize)
if err != nil {
response.HandleError(c, err)
return
}
response.Success(c, gin.H{
"items": notifications,
"total": total,
"page": page,
})
}
func (h *Handler) MarkAsRead(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, response.NewBadRequest("invalid id"))
return
}
tenantID := c.GetString("tenant_id")
if err := h.svc.MarkAsRead(c.Request.Context(), tenantID, uint(id)); err != nil {
response.HandleError(c, err)
return
}
response.Success(c, nil)
}
func (h *Handler) MarkAllAsRead(c *gin.Context) {
userID := c.GetUint("user_id")
tenantID := c.GetString("tenant_id")
if err := h.svc.MarkAllAsRead(c.Request.Context(), tenantID, userID); err != nil {
response.HandleError(c, err)
return
}
response.Success(c, nil)
}
func (h *Handler) CountUnread(c *gin.Context) {
userID := c.GetUint("user_id")
tenantID := c.GetString("tenant_id")
count, err := h.svc.CountUnread(c.Request.Context(), tenantID, userID)
if err != nil {
response.HandleError(c, err)
return
}
response.Success(c, gin.H{"count": count})
}
func (h *Handler) Delete(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, response.NewBadRequest("invalid id"))
return
}
tenantID := c.GetString("tenant_id")
if err := h.svc.Delete(c.Request.Context(), tenantID, uint(id)); err != nil {
response.HandleError(c, err)
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "deleted"})
}

View File

@ -0,0 +1,43 @@
package interfaces
import (
"mengstack/internal/app/middleware"
"mengstack/internal/modules/notification/application"
"mengstack/internal/modules/notification/infrastructure"
"github.com/gin-gonic/gin"
"go.uber.org/fx"
"go.uber.org/zap"
"gorm.io/gorm"
)
func SetupRoutes(
r *gin.Engine,
h *Handler,
authMW gin.HandlerFunc,
) {
notifications := r.Group("/api/v1/notifications")
notifications.Use(authMW, middleware.MultiTenant())
{
notifications.POST("", h.Create)
notifications.GET("/:id", h.Get)
notifications.DELETE("/:id", h.Delete)
notifications.PUT("/:id/read", h.MarkAsRead)
notifications.PUT("/read-all", h.MarkAllAsRead)
notifications.GET("/unread-count", h.CountUnread)
notifications.GET("/user/:userId", h.ListByUser)
}
}
var Module = fx.Module("notification",
fx.Provide(
infrastructure.NewNotificationRepository,
application.NewService,
NewHandler,
),
fx.Invoke(func(db *gorm.DB, log *zap.Logger) {
if err := infrastructure.Migrate(db); err != nil {
log.Fatal("notification migration failed", zap.Error(err))
}
}),
)