package shareser import ( "context" "errors" "sync" "testing" "time" "go.mongodb.org/mongo-driver/bson/primitive" ) type fakeRecommendShareDeduper struct { mu sync.Mutex keys map[string]bool ttls []time.Duration err error } func (f *fakeRecommendShareDeduper) SetNXContext( _ context.Context, key string, _ interface{}, expiration time.Duration, ) (bool, error) { f.mu.Lock() defer f.mu.Unlock() if f.err != nil { return false, f.err } f.ttls = append(f.ttls, expiration) if f.keys == nil { f.keys = make(map[string]bool) } if f.keys[key] { return false, nil } f.keys[key] = true return true, nil } func (f *fakeRecommendShareDeduper) DelContext( _ context.Context, keys ...string, ) (int64, error) { f.mu.Lock() defer f.mu.Unlock() var removed int64 for _, key := range keys { if f.keys[key] { delete(f.keys, key) removed++ } } return removed, nil } func TestIncrementRecommendShareOnceDeduplicatesEvent(t *testing.T) { store := &fakeRecommendShareDeduper{} id := primitive.NewObjectID() now := time.Date(2026, 7, 31, 8, 0, 0, 0, time.UTC) calls := 0 increment := func(context.Context, primitive.ObjectID) error { calls++ return nil } for i := 0; i < 2; i++ { if err := incrementRecommendShareOnceWith( context.Background(), store, 123, id, "event-1", now, increment, ); err != nil { t.Fatal(err) } } if calls != 1 { t.Fatalf("increment calls = %d, want 1", calls) } if len(store.ttls) != 3 || store.ttls[0] != recommendShareDailyDedupTTL || store.ttls[1] != recommendShareEventDedupTTL || store.ttls[2] != recommendShareDailyDedupTTL { t.Fatalf("dedupe TTL calls = %v", store.ttls) } } func TestIncrementRecommendShareOnceCapsDifferentEventsPerDay(t *testing.T) { store := &fakeRecommendShareDeduper{} id := primitive.NewObjectID() now := time.Date(2026, 7, 31, 8, 0, 0, 0, time.UTC) calls := 0 increment := func(context.Context, primitive.ObjectID) error { calls++ return nil } for _, eventID := range []string{"event-1", "event-2", "event-3"} { if err := incrementRecommendShareOnceWith( context.Background(), store, 123, id, eventID, now, increment, ); err != nil { t.Fatal(err) } } if calls != 1 { t.Fatalf("increment calls = %d, want one per user/video/day", calls) } } func TestIncrementRecommendShareOnceEventRetryDoesNotConsumeNextDayQuota(t *testing.T) { store := &fakeRecommendShareDeduper{} id := primitive.NewObjectID() first := time.Date(2026, 7, 31, 15, 59, 0, 0, time.UTC) nextDay := first.Add(2 * time.Minute) calls := 0 increment := func(context.Context, primitive.ObjectID) error { calls++ return nil } for _, step := range []struct { eventID string now time.Time }{ {eventID: "event-1", now: first}, {eventID: "event-1", now: nextDay}, {eventID: "event-2", now: nextDay}, } { if err := incrementRecommendShareOnceWith( context.Background(), store, 123, id, step.eventID, step.now, increment, ); err != nil { t.Fatal(err) } } if calls != 2 { t.Fatalf("increment calls = %d, want one on each day", calls) } } func TestIncrementRecommendShareOnceLegacyScopesByCSTDay(t *testing.T) { store := &fakeRecommendShareDeduper{} id := primitive.NewObjectID() incremented := 0 increment := func(context.Context, primitive.ObjectID) error { incremented++ return nil } first := time.Date(2026, 7, 31, 15, 59, 0, 0, time.UTC) second := first.Add(2 * time.Minute) for _, now := range []time.Time{first, first, second} { if err := incrementRecommendShareOnceWith( context.Background(), store, 123, id, "", now, increment, ); err != nil { t.Fatal(err) } } if incremented != 2 { t.Fatalf("increment calls = %d, want 2 days", incremented) } } func TestIncrementRecommendShareOnceRedisFailureDoesNotIncrement(t *testing.T) { wantErr := errors.New("redis unavailable") store := &fakeRecommendShareDeduper{err: wantErr} calls := 0 err := incrementRecommendShareOnceWith( context.Background(), store, 123, primitive.NewObjectID(), "event", time.Now(), func(context.Context, primitive.ObjectID) error { calls++ return nil }, ) if !errors.Is(err, wantErr) { t.Fatalf("error = %v, want %v", err, wantErr) } if calls != 0 { t.Fatalf("increment calls = %d, want 0", calls) } } func TestIncrementRecommendShareOnceReleasesKeyAfterFailure(t *testing.T) { store := &fakeRecommendShareDeduper{} id := primitive.NewObjectID() wantErr := errors.New("injected") calls := 0 increment := func(context.Context, primitive.ObjectID) error { calls++ if calls == 1 { return wantErr } return nil } if err := incrementRecommendShareOnceWith( context.Background(), store, 123, id, "event", time.Now(), increment, ); !errors.Is(err, wantErr) { t.Fatalf("first error = %v, want injected", err) } if err := incrementRecommendShareOnceWith( context.Background(), store, 123, id, "event", time.Now(), increment, ); err != nil { t.Fatalf("retry error = %v", err) } if calls != 2 { t.Fatalf("increment calls = %d, want 2", calls) } } func TestIncrementRecommendShareOnceConcurrent(t *testing.T) { store := &fakeRecommendShareDeduper{} id := primitive.NewObjectID() var mu sync.Mutex calls := 0 increment := func(context.Context, primitive.ObjectID) error { mu.Lock() calls++ mu.Unlock() return nil } var wg sync.WaitGroup for i := 0; i < 32; i++ { wg.Add(1) go func() { defer wg.Done() if err := incrementRecommendShareOnceWith( context.Background(), store, 123, id, "event", time.Now(), increment, ); err != nil { t.Errorf("increment error = %v", err) } }() } wg.Wait() if calls != 1 { t.Fatalf("increment calls = %d, want 1", calls) } } func TestIncrementRecommendShareOnceRejectsLongEventID(t *testing.T) { err := incrementRecommendShareOnceWith( context.Background(), &fakeRecommendShareDeduper{}, 123, primitive.NewObjectID(), string(make([]byte, recommendShareEventIDMaxLength+1)), time.Now(), func(context.Context, primitive.ObjectID) error { return nil }, ) if err == nil { t.Fatal("expected long eventId error") } }