Skip to main content

photon_backend/consumer_group/
coordinator.rs

1//! Consumer group coordinator trait and implementations.
2
3use async_trait::async_trait;
4
5use crate::consumer_group::static_assignment;
6use crate::error::Result;
7
8/// Registered group member.
9#[derive(Debug, Clone)]
10pub struct GroupMember {
11    /// Consumer group id.
12    pub group_id: String,
13    /// Unique member instance id within the group.
14    pub instance_id: String,
15    /// Topic this member consumes.
16    pub topic_name: String,
17    /// Virtual shard count for the topic.
18    pub shard_count: u32,
19}
20
21/// Assigns virtual shards to group members.
22#[async_trait]
23pub trait ConsumerGroupCoordinator: Send + Sync {
24    /// Register member and return assigned shard ids.
25    async fn register(&self, member: GroupMember) -> Result<Vec<u32>>;
26
27    /// Heartbeat to keep assignment alive (fleet lease store).
28    async fn heartbeat(&self, group_id: &str, instance_id: &str) -> Result<()> {
29        let _ = (group_id, instance_id);
30        Ok(())
31    }
32
33    /// Current assignment for a member.
34    async fn assigned_shards(&self, group_id: &str, instance_id: &str) -> Result<Vec<u32>>;
35}
36
37/// Env-based static assignment (`PHOTON_GROUP_SHARD_ASSIGNMENT`).
38pub struct StaticGroupCoordinator;
39
40#[async_trait]
41impl ConsumerGroupCoordinator for StaticGroupCoordinator {
42    async fn register(&self, member: GroupMember) -> Result<Vec<u32>> {
43        Ok(static_assignment::static_assigned_shards(
44            member.shard_count,
45        ))
46    }
47
48    async fn assigned_shards(&self, group_id: &str, instance_id: &str) -> Result<Vec<u32>> {
49        let _ = (group_id, instance_id);
50        Ok(static_assignment::static_assigned_shards(
51            static_assignment::shard_count_from_env().unwrap_or(32),
52        ))
53    }
54}
55
56/// Fleet coordinator backed by a [`super::lease_store::LeaseStore`].
57pub struct FleetGroupCoordinator<L: super::lease_store::LeaseStore> {
58    store: L,
59    lease_ttl_secs: u64,
60}
61
62impl<L: super::lease_store::LeaseStore> FleetGroupCoordinator<L> {
63    /// Create a coordinator with the given lease store and lease TTL.
64    pub const fn new(store: L, lease_ttl_secs: u64) -> Self {
65        Self {
66            store,
67            lease_ttl_secs,
68        }
69    }
70
71    fn range_assign(member_index: u32, member_count: u32, shard_count: u32) -> Vec<u32> {
72        if member_count == 0 {
73            return Vec::new();
74        }
75        let per = shard_count.div_ceil(member_count);
76        let start = member_index * per;
77        let end = (start + per).min(shard_count);
78        (start..end).collect()
79    }
80}
81
82#[async_trait]
83impl<L: super::lease_store::LeaseStore> ConsumerGroupCoordinator for FleetGroupCoordinator<L> {
84    async fn register(&self, member: GroupMember) -> Result<Vec<u32>> {
85        let member_index = member.instance_id.parse::<u32>().unwrap_or_else(|_| {
86            u32::try_from(member.instance_id.len()).unwrap_or(0) % member.shard_count.max(1)
87        });
88        let member_count = static_assignment::member_count_from_env().unwrap_or(2);
89        let shards = Self::range_assign(member_index, member_count, member.shard_count);
90        for shard_id in &shards {
91            self.store
92                .claim(super::lease_store::ConsumerLease {
93                    group_id: member.group_id.clone(),
94                    shard_id: *shard_id,
95                    instance_id: member.instance_id.clone(),
96                    ttl_secs: self.lease_ttl_secs,
97                })
98                .await?;
99        }
100        Ok(shards)
101    }
102
103    async fn heartbeat(&self, group_id: &str, instance_id: &str) -> Result<()> {
104        self.store
105            .renew(group_id, instance_id, self.lease_ttl_secs)
106            .await
107    }
108
109    async fn assigned_shards(&self, group_id: &str, instance_id: &str) -> Result<Vec<u32>> {
110        self.store.list_for_instance(group_id, instance_id).await
111    }
112}