use crate::Identifier;
use std::collections::HashMap;
use std::sync::Mutex;
#[derive(Debug, Default, Clone)]
struct GroupAssignment {
partitions: Vec<u32>,
generation: u64,
cursor: usize,
}
#[derive(Debug, Default)]
pub struct ConsumerGroupClientState {
assignments: Mutex<HashMap<String, GroupAssignment>>,
balanced_cursors: Mutex<HashMap<String, usize>>,
partition_counts: Mutex<HashMap<String, u32>>,
joined_groups: Mutex<HashMap<String, (Identifier, Identifier, Identifier)>>,
}
impl ConsumerGroupClientState {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn has_assignment(&self, key: &str) -> bool {
self.assignments
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(key)
.is_some_and(|assignment| !assignment.partitions.is_empty())
}
pub fn set_assignment(&self, key: String, generation: u64, partitions: Vec<u32>) {
let mut map = self
.assignments
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let entry = map.entry(key).or_default();
if entry.generation != generation {
entry.cursor = 0;
}
entry.generation = generation;
entry.partitions = partitions;
}
pub fn invalidate_assignment(&self, key: &str) {
self.assignments
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(key);
}
#[must_use]
pub fn next_group_partition(&self, key: &str) -> Option<u32> {
let mut map = self
.assignments
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let assignment = map.get_mut(key)?;
if assignment.partitions.is_empty() {
return None;
}
let index = assignment.cursor % assignment.partitions.len();
assignment.cursor = assignment.cursor.wrapping_add(1);
Some(assignment.partitions[index])
}
#[must_use]
#[allow(clippy::cast_possible_truncation)]
pub fn next_balanced_partition(&self, key: &str, partition_count: u32) -> u32 {
if partition_count == 0 {
return 0;
}
let mut map = self
.balanced_cursors
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let cursor = map.entry(key.to_owned()).or_default();
let partition = (*cursor % partition_count as usize) as u32;
*cursor = cursor.wrapping_add(1);
partition
}
#[must_use]
pub fn partition_count(&self, key: &str) -> Option<u32> {
self.partition_counts
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(key)
.copied()
}
pub fn set_partition_count(&self, key: String, partition_count: u32) {
self.partition_counts
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(key, partition_count);
}
pub fn register_group(
&self,
key: String,
stream_id: Identifier,
topic_id: Identifier,
group_id: Identifier,
) {
self.joined_groups
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(key, (stream_id, topic_id, group_id));
}
pub fn deregister_group(&self, key: &str) {
self.joined_groups
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(key);
}
#[must_use]
pub fn is_registered(&self, key: &str) -> bool {
self.joined_groups
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.contains_key(key)
}
#[must_use]
pub fn registered_groups(&self) -> Vec<(Identifier, Identifier, Identifier)> {
self.joined_groups
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.values()
.cloned()
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn group_partition_round_robins_then_wraps() {
let state = ConsumerGroupClientState::new();
state.set_assignment("s|t|g".to_owned(), 1, vec![0, 1, 2]);
let picks: Vec<u32> = (0..4)
.map(|_| state.next_group_partition("s|t|g").unwrap())
.collect();
assert_eq!(picks, vec![0, 1, 2, 0]);
}
#[test]
fn generation_change_resets_cursor() {
let state = ConsumerGroupClientState::new();
state.set_assignment("s|t|g".to_owned(), 1, vec![0, 1, 2]);
assert_eq!(state.next_group_partition("s|t|g"), Some(0));
assert_eq!(state.next_group_partition("s|t|g"), Some(1));
state.set_assignment("s|t|g".to_owned(), 2, vec![5]);
assert_eq!(state.next_group_partition("s|t|g"), Some(5));
}
#[test]
fn balanced_round_robins() {
let state = ConsumerGroupClientState::new();
let picks: Vec<u32> = (0..4)
.map(|_| state.next_balanced_partition("s|t", 3))
.collect();
assert_eq!(picks, vec![0, 1, 2, 0]);
}
#[test]
fn missing_assignment_yields_none() {
let state = ConsumerGroupClientState::new();
assert!(!state.has_assignment("s|t|g"));
assert_eq!(state.next_group_partition("s|t|g"), None);
}
#[test]
fn member_holding_no_partitions_stays_registered() {
let state = ConsumerGroupClientState::new();
let id = Identifier::named("g").unwrap();
state.register_group("s|t|g".to_owned(), id.clone(), id.clone(), id);
state.set_assignment("s|t|g".to_owned(), 1, Vec::new());
assert!(!state.has_assignment("s|t|g"));
assert!(state.is_registered("s|t|g"));
state.deregister_group("s|t|g");
assert!(!state.is_registered("s|t|g"));
}
}