use std::cmp::Ordering;
use std::collections::hash_map::Entry;
use ahash::{AHashMap, AHashSet};
use itertools::Itertools;
use crate::segment::data_types::groups::GroupId;
use crate::segment::json_path::JsonPath;
use crate::segment::spaces::tools::{peek_top_largest_iterable, peek_top_smallest_iterable};
use crate::segment::types::{ExtendedPointId, Order, PayloadContainer, PointIdType, ScoredPoint};
use serde_json::Value;
use super::{AggregatorError, Group};
const LARGEST_REASONABLE_ALLOCATION_SIZE: usize = 1_048_576;
type Hits = AHashMap<PointIdType, ScoredPoint>;
pub struct GroupsAggregator {
groups: AHashMap<GroupId, Hits>,
max_group_size: usize,
grouped_by: JsonPath,
max_groups: usize,
full_groups: AHashSet<GroupId>,
group_best_scores: AHashMap<GroupId, ScoredPoint>,
all_ids: AHashSet<ExtendedPointId>,
order: Option<Order>,
}
impl GroupsAggregator {
pub fn new(
groups: usize,
group_size: usize,
grouped_by: JsonPath,
order: Option<Order>,
) -> Self {
Self {
groups: AHashMap::with_capacity(groups.min(LARGEST_REASONABLE_ALLOCATION_SIZE)),
max_group_size: group_size,
grouped_by,
max_groups: groups,
full_groups: AHashSet::with_capacity(groups.min(LARGEST_REASONABLE_ALLOCATION_SIZE)),
group_best_scores: AHashMap::with_capacity(
groups.min(LARGEST_REASONABLE_ALLOCATION_SIZE),
),
all_ids: AHashSet::with_capacity(
groups
.saturating_mul(group_size)
.min(LARGEST_REASONABLE_ALLOCATION_SIZE),
),
order,
}
}
fn add_point(&mut self, point: &ScoredPoint) -> Result<(), AggregatorError> {
let payload_values: Vec<_> = point
.payload
.as_ref()
.map(|p| {
p.get_value(&self.grouped_by)
.into_iter()
.flat_map(|v| match v {
Value::Array(arr) => arr.iter().collect(),
Value::Null
| Value::Bool(_)
| Value::Number(_)
| Value::String(_)
| Value::Object(_) => vec![v],
})
.collect()
})
.ok_or(AggregatorError::KeyNotFound)?;
let unique_group_keys: Vec<_> =
itertools::process_results(payload_values.into_iter().map(GroupId::try_from), |iter| {
iter.unique().collect()
})
.map_err(|_| AggregatorError::BadKeyType)?;
for group_key in unique_group_keys {
let group = self.groups.entry(group_key.clone()).or_insert_with(|| {
AHashMap::with_capacity(self.max_group_size.min(LARGEST_REASONABLE_ALLOCATION_SIZE))
});
let entry = group.entry(point.id);
match entry {
Entry::Occupied(mut o) => {
if o.get().version < point.version {
o.insert(point.clone());
}
}
Entry::Vacant(v) => {
v.insert(point.clone());
self.all_ids.insert(point.id);
}
}
if group.len() == self.max_group_size {
self.full_groups.insert(group_key.clone());
}
self.group_best_scores
.entry(group_key.clone())
.and_modify(|other_score| {
let ordering = match self.order {
Some(Order::LargeBetter) => point.cmp(other_score),
Some(Order::SmallBetter) => (*other_score).cmp(point),
None => Ordering::Equal, };
if ordering == Ordering::Greater {
*other_score = point.clone();
}
})
.or_insert(point.clone());
}
Ok(())
}
pub fn add_points(&mut self, points: &[ScoredPoint]) {
for point in points {
match self.add_point(point) {
Ok(()) | Err(AggregatorError::KeyNotFound | AggregatorError::BadKeyType) => {
}
}
}
}
#[cfg(test)]
pub(super) fn len(&self) -> usize {
self.groups.len()
}
fn best_group_keys(&self) -> Vec<GroupId> {
let mut pairs: Vec<_> = self.group_best_scores.iter().collect();
pairs.sort_unstable_by(|(_, score1), (_, score2)| match self.order {
Some(Order::LargeBetter) => score2.cmp(score1),
Some(Order::SmallBetter) => score1.cmp(score2),
None => Ordering::Equal,
});
pairs
.iter()
.take(self.max_groups)
.map(|(k, _)| (*k).clone())
.collect()
}
pub fn keys_of_unfilled_best_groups(&self) -> Vec<Value> {
let best_group_keys: AHashSet<_> = self.best_group_keys().into_iter().collect();
best_group_keys
.difference(&self.full_groups)
.cloned()
.map_into()
.collect()
}
pub fn keys_of_filled_groups(&self) -> Vec<Value> {
self.full_groups.iter().cloned().map_into().collect()
}
pub fn len_of_filled_best_groups(&self) -> usize {
let best_group_keys: AHashSet<_> = self.best_group_keys().into_iter().collect();
best_group_keys.intersection(&self.full_groups).count()
}
pub fn ids(&self) -> &AHashSet<ExtendedPointId> {
&self.all_ids
}
pub fn distill(mut self) -> Vec<Group> {
let best_groups = self.best_group_keys();
let mut groups = Vec::with_capacity(best_groups.len());
for group_key in best_groups {
let mut group = self.groups.remove(&group_key).unwrap();
let scored_points_iter = group.drain().map(|(_, hit)| hit);
let hits = match self.order {
Some(Order::LargeBetter) => {
peek_top_largest_iterable(scored_points_iter, self.max_group_size)
}
Some(Order::SmallBetter) => {
peek_top_smallest_iterable(scored_points_iter, self.max_group_size)
}
None => scored_points_iter.take(self.max_group_size).collect(),
};
groups.push(Group {
hits,
key: group_key,
});
}
groups
}
}
#[cfg(test)]
mod unit_tests {
use crate::common::types::ScoreType;
use crate::segment::payload_json;
use serde_json::json;
use super::*;
fn point(idx: u64, score: ScoreType, payloads: Value) -> ScoredPoint {
ScoredPoint {
id: idx.into(),
version: 0,
score,
payload: Some(payload_json! { "docId": payloads }),
vector: None,
shard_key: None,
order_value: None,
}
}
fn empty_point(idx: u64, score: ScoreType) -> ScoredPoint {
ScoredPoint {
id: idx.into(),
version: 0,
score,
payload: None,
vector: None,
shard_key: None,
order_value: None,
}
}
#[test]
fn test_group_with_multiple_payload_values() {
let scored_points = vec![
point(1, 0.99, json!(["a", "a"])),
point(2, 0.85, json!(["a", "b"])),
point(3, 0.75, json!("b")),
];
let mut aggregator =
GroupsAggregator::new(3, 2, "docId".parse().unwrap(), Some(Order::LargeBetter));
for point in &scored_points {
aggregator.add_point(point).unwrap();
}
let result = aggregator.distill();
assert_eq!(result.len(), 2);
assert_eq!(result[0].hits.len(), 2);
assert_eq!(result[0].hits[0].id, 1.into());
assert_eq!(result[0].hits[1].id, 2.into());
assert_eq!(result[1].hits.len(), 2);
assert_eq!(result[1].hits[0].id, 2.into());
assert_eq!(result[1].hits[1].id, 3.into());
}
struct Case {
point: ScoredPoint,
key: Value,
group_size: usize,
groups_count: usize,
expected_result: Result<(), AggregatorError>,
}
impl Case {
fn new(
key: Value,
group_size: usize,
groups_count: usize,
expected_result: Result<(), AggregatorError>,
point: ScoredPoint,
) -> Self {
Self {
point,
key,
group_size,
groups_count,
expected_result,
}
}
}
#[test]
fn it_adds_single_points() {
let mut aggregator =
GroupsAggregator::new(4, 3, "docId".parse().unwrap(), Some(Order::LargeBetter));
#[rustfmt::skip]
[
Case::new(json!("a"), 1, 1, Ok(()), point(1, 0.99, json!("a"))),
Case::new(json!("a"), 1, 1, Ok(()), point(1, 0.97, json!("a"))), Case::new(json!("a"), 2, 2, Ok(()), point(2, 0.81, json!(["a", "b"]))), Case::new(json!("b"), 2, 2, Ok(()), point(3, 0.84, json!("b"))), Case::new(json!("a"), 3, 2, Ok(()), point(4, 0.9, json!("a"))), Case::new(json!(3), 1, 3, Ok(()), point(5, 0.4, json!(3))), Case::new(json!("d"), 1, 4, Ok(()), point(6, 0.3, json!("d"))),
Case::new(json!("a"), 4, 4, Ok(()), point(100, 0.31, json!("a"))), Case::new(json!("a"), 5, 4, Ok(()), point(101, 0.32, json!("a"))), Case::new(json!("a"), 6, 4, Ok(()), point(102, 0.33, json!("a"))), Case::new(json!("a"), 7, 4, Ok(()), point(103, 0.34, json!("a"))), Case::new(json!("a"), 8, 4, Ok(()), point(104, 0.35, json!("a"))), Case::new(json!("a"), 9, 4, Ok(()), point(105, 0.36, json!("a"))), Case::new(json!("b"), 3, 4, Ok(()), point(7, 1.0, json!("b"))),
Case::new(json!("false"), 0, 4, Err(AggregatorError::BadKeyType), point(8, 1.0, json!(false))),
Case::new(json!("none"), 0, 4, Err(AggregatorError::KeyNotFound), empty_point(9, 1.0)),
Case::new(json!(3), 2, 4, Ok(()), point(10, 0.6, json!(3))),
Case::new(json!(3), 3, 4, Ok(()), point(11, 0.1, json!(3))),
]
.into_iter()
.enumerate()
.for_each(|(case_idx, case)| {
let result = aggregator.add_point(&case.point);
assert_eq!(result, case.expected_result, "case {case_idx}");
assert_eq!(aggregator.len(), case.groups_count, "case {case_idx}");
let key = &GroupId::try_from(&case.key).unwrap();
if case.group_size > 0 {
assert_eq!(
aggregator.groups.get(key).unwrap().len(),
case.group_size,
"case {case_idx}"
);
} else {
assert!(!aggregator.groups.contains_key(key), "case {case_idx}");
}
});
assert_eq!(aggregator.full_groups.len(), 3);
assert_eq!(aggregator.keys_of_unfilled_best_groups(), vec![json!("d")]);
assert_eq!(aggregator.len_of_filled_best_groups(), 3);
let groups = aggregator.distill();
#[rustfmt::skip]
let expected_groups = vec![
(
GroupId::from("b"),
vec![
empty_point(7, 1.0),
empty_point(3, 0.84),
empty_point(2, 0.81),
],
),
(
GroupId::from("a"),
vec![
empty_point(1, 0.99),
empty_point(4, 0.9),
empty_point(2, 0.81)
],
),
(
GroupId::try_from(&json!(3)).unwrap(),
vec![
empty_point(10, 0.6),
empty_point(5, 0.4),
empty_point(11, 0.1),
],
),
(
GroupId::from("d"),
vec![
empty_point(6, 0.3),
],
),
];
for ((expected_key, expected_group_points), group) in
expected_groups.into_iter().zip(groups)
{
assert_eq!(expected_key, group.key);
let expected_id_score: Vec<_> = expected_group_points
.into_iter()
.map(|x| (x.id, x.score))
.collect();
let group_id_score: Vec<_> = group.hits.into_iter().map(|x| (x.id, x.score)).collect();
assert_eq!(expected_id_score, group_id_score);
}
}
#[test]
fn test_aggregate_less_groups() {
let mut aggregator =
GroupsAggregator::new(3, 2, "docId".parse().unwrap(), Some(Order::LargeBetter));
[
point(1, 0.99, json!("a")),
point(1, 0.97, json!("a")), point(2, 0.81, json!(["a", "b"])), point(3, 0.84, json!("b")), point(4, 0.9, json!("a")), point(5, 0.4, json!(3)), point(6, 0.3, json!("d")),
point(100, 0.31, json!("a")), point(101, 0.32, json!("a")), point(102, 0.33, json!("a")), point(103, 0.34, json!("a")), point(104, 0.35, json!("a")), point(105, 0.36, json!("a")), point(7, 1.0, json!("b")),
point(10, 0.6, json!(3)),
point(11, 0.1, json!(3)),
]
.iter()
.for_each(|point| {
aggregator.add_point(point).unwrap();
});
let groups = aggregator.distill();
#[rustfmt::skip]
let expected_groups = vec![
(
GroupId::from("b"),
vec![
empty_point(7, 1.0),
empty_point(3, 0.84),
],
),
(
GroupId::from("a"),
vec![
empty_point(1, 0.99),
empty_point(4, 0.9),
],
),
(
GroupId::try_from(&json!(3)).unwrap(),
vec![
empty_point(10, 0.6),
empty_point(5, 0.4),
],
),
];
for ((key, expected_group_points), group) in expected_groups.into_iter().zip(groups) {
assert_eq!(key, group.key);
let expected_id_score: Vec<_> = expected_group_points
.into_iter()
.map(|x| (x.id, x.score))
.collect();
let group_id_score: Vec<_> = group.hits.into_iter().map(|x| (x.id, x.score)).collect();
assert_eq!(expected_id_score, group_id_score);
}
}
#[test]
fn test_large_groups_and_group_size_do_not_panic() {
let mut aggregator = GroupsAggregator::new(
usize::MAX,
usize::MAX,
"docId".parse().unwrap(),
Some(Order::LargeBetter),
);
aggregator.add_point(&point(1, 0.5, json!("a"))).unwrap();
}
}