mengstack-api/internal/modules/notification/application/service_test.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

280 lines
6.8 KiB
Go

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")
}
}