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