diff --git a/app/models/push/subscription.rb b/app/models/push/subscription.rb index bd78482..e4345e1 100644 --- a/app/models/push/subscription.rb +++ b/app/models/push/subscription.rb @@ -14,9 +14,9 @@ class Push::Subscription < ApplicationRecord validates :endpoint, presence: true validate :validate_endpoint_url - def notification(**params) + def notification(badge: user.memberships.unread.count, **params) # Defer DNS lookup to the delivery worker to prevent rebinding - WebPush::Notification.new(**params, badge: user.memberships.unread.count, endpoint: endpoint, endpoint_ip_resolver: method(:resolved_endpoint_ip), p256dh_key: p256dh_key, auth_key: auth_key) + WebPush::Notification.new(**params, badge: badge, endpoint: endpoint, endpoint_ip_resolver: method(:resolved_endpoint_ip), p256dh_key: p256dh_key, auth_key: auth_key) end # Validate at point of use, not just when saved. diff --git a/lib/web_push/pool.rb b/lib/web_push/pool.rb index aeccea5..1326d3d 100644 --- a/lib/web_push/pool.rb +++ b/lib/web_push/pool.rb @@ -10,8 +10,12 @@ class WebPush::Pool end def queue(payload, subscriptions) - subscriptions.find_each do |subscription| - deliver_later(payload, subscription) + subscriptions.find_in_batches do |batch| + unread_counts = Membership.unread.where(user_id: batch.map(&:user_id)).group(:user_id).count + + batch.each do |subscription| + deliver_later(payload, subscription, badge: unread_counts.fetch(subscription.user_id, 0)) + end end end @@ -22,9 +26,9 @@ class WebPush::Pool end private - def deliver_later(payload, subscription) + def deliver_later(payload, subscription, badge:) # Ensure any AR operations happen before we post to the thread pool - notification = subscription.notification(**payload) + notification = subscription.notification(**payload, badge: badge) subscription_id = subscription.id delivery_pool.post do diff --git a/test/lib/web_push/pool_test.rb b/test/lib/web_push/pool_test.rb new file mode 100644 index 0000000..1377cdb --- /dev/null +++ b/test/lib/web_push/pool_test.rb @@ -0,0 +1,53 @@ +require "test_helper" + +class WebPush::PoolTest < ActiveSupport::TestCase + setup do + stub_web_push_dns_resolution + @pool = Rails.configuration.x.web_push_pool + @payload = { title: "Designers", body: "Hello", path: Rails.application.routes.url_helpers.room_path(rooms(:designers)) } + end + + test "every notification carries its subscriber's unread room count as the badge" do + memberships(:jason_pets).update! unread_at: Time.current + memberships(:kevin_hq).update! unread_at: Time.current + memberships(:kevin_david_and_kevin).update! unread_at: Time.current + subscriptions = Push::Subscription.where(user: users(:david, :jason, :kevin)) + expected = subscriptions.to_h { |subscription| [ subscription.endpoint, subscription.user.memberships.unread.count ] } + assert_equal [ 0, 1, 2 ], expected.values.sort + + badges = Concurrent::Hash.new + WebPush.stubs(:payload_send).with { |options| badges[options[:endpoint]] = JSON.parse(options[:message]).dig("options", "data", "badge") } + @pool.queue(@payload, subscriptions) + wait_for_deliveries(3) + + assert_equal expected, badges + end + + test "queueing counts unread rooms per batch, not per subscription" do + WebPush.stubs(:payload_send) + + two = Push::Subscription.where(user: users(:david, :jason)) + four = Push::Subscription.where(user: users(:david, :jason, :jz, :kevin)) + assert_equal [ 2, 4 ], [ two.count, four.count ] + + queries_for_two = count_queries { @pool.queue(@payload, two) } + queries_for_four = count_queries { @pool.queue(@payload, four) } + wait_for_deliveries(6) + + assert_equal queries_for_two, queries_for_four + end + + private + def wait_for_deliveries(count) + Timeout.timeout(2) { sleep 0.01 while @pool.delivery_pool.completed_task_count < count } + end + + def count_queries(&block) + count = 0 + counter = ->(*, payload) { count += 1 unless payload[:name] == "SCHEMA" || payload[:cached] } + ActiveRecord::Base.uncached do + ActiveSupport::Notifications.subscribed(counter, "sql.active_record", &block) + end + count + end +end diff --git a/test/models/push/subscription_test.rb b/test/models/push/subscription_test.rb index 6dce213..f6df6cd 100644 --- a/test/models/push/subscription_test.rb +++ b/test/models/push/subscription_test.rb @@ -125,6 +125,15 @@ class Push::SubscriptionTest < ActiveSupport::TestCase subscription.notification(title: "t", body: "b", path: "/").deliver end + test "the badge defaults to the subscriber's unread room count" do + memberships(:david_pets).update! unread_at: Time.current + memberships(:david_hq).update! unread_at: Time.current + subscription = build_subscription(endpoint: "https://fcm.googleapis.com/fcm/send/abc123") + + WebPush.expects(:payload_send).with { |options| JSON.parse(options[:message]).dig("options", "data", "badge") == 2 } + subscription.notification(title: "t", body: "b", path: "/").deliver + end + test "delivery sends with the pinned endpoint_ip" do subscription = build_subscription(endpoint: "https://fcm.googleapis.com/fcm/send/abc123")