use crate::segment::json_path::JsonPath;
use crate::segment::types::{Condition, Filter, Order, ScoredPoint};
use super::{Group, GroupsAggregator, except_on, match_on, merge_filter, shape_candidates_query};
use crate::shard::query::ShardQueryRequest;
const MAX_GET_GROUPS_REQUESTS: usize = 5;
const MAX_GROUP_FILLING_REQUESTS: usize = 5;
#[derive(Copy, Clone, Debug)]
pub struct RequestBudget {
pub collect: usize,
pub fill: usize,
}
impl Default for RequestBudget {
fn default() -> Self {
Self {
collect: MAX_GET_GROUPS_REQUESTS,
fill: MAX_GROUP_FILLING_REQUESTS,
}
}
}
pub struct GroupByDriver {
base_query: ShardQueryRequest,
group_by: JsonPath,
groups: usize,
group_size: usize,
candidates_limit: usize,
aggregator: GroupsAggregator,
state: State,
}
#[derive(Copy, Clone)]
enum State {
Collecting {
requests_left: usize,
fill_budget: usize,
},
Filling {
requests_left: usize,
},
Done,
}
impl GroupByDriver {
pub fn new(
base_query: ShardQueryRequest,
group_by: JsonPath,
groups: usize,
group_size: usize,
order: Option<Order>,
budget: RequestBudget,
) -> Self {
let state = if groups == 0 || group_size == 0 {
State::Done
} else {
State::Collecting {
requests_left: budget.collect,
fill_budget: budget.fill,
}
};
Self {
aggregator: GroupsAggregator::new(groups, group_size, group_by.clone(), order),
base_query,
group_by,
groups,
group_size,
candidates_limit: groups.saturating_mul(group_size),
state,
}
}
pub fn next_request(&mut self) -> Option<ShardQueryRequest> {
loop {
match &mut self.state {
State::Collecting {
requests_left,
fill_budget,
} => {
if *requests_left == 0 {
self.state = State::Filling {
requests_left: *fill_budget,
};
continue;
}
*requests_left -= 1;
let mut request = self.base_query.clone();
let full_groups = self.aggregator.keys_of_filled_groups();
if !full_groups.is_empty() {
let except_any = except_on(&self.group_by, &full_groups);
if !except_any.is_empty() {
merge_filter(
&mut request.filter,
Filter {
must: Some(except_any),
..Default::default()
},
);
}
}
self.exclude_aggregated_points(&mut request);
shape_candidates_query(
&mut request,
&self.group_by,
self.candidates_limit,
self.group_size,
);
return Some(request);
}
State::Filling { requests_left } => {
if *requests_left == 0 {
self.state = State::Done;
continue;
}
*requests_left -= 1;
let mut request = self.base_query.clone();
let unsatisfied_groups = self.aggregator.keys_of_unfilled_best_groups();
let match_any = match_on(&self.group_by, &unsatisfied_groups);
if !match_any.is_empty() {
merge_filter(
&mut request.filter,
Filter {
must: Some(match_any),
..Default::default()
},
);
}
self.exclude_aggregated_points(&mut request);
shape_candidates_query(
&mut request,
&self.group_by,
self.candidates_limit,
self.group_size,
);
return Some(request);
}
State::Done => return None,
}
}
}
pub fn add_points(&mut self, points: &[ScoredPoint]) {
self.aggregator.add_points(points);
let enough_groups = self.aggregator.len_of_filled_best_groups() >= self.groups;
match self.state {
State::Collecting { fill_budget, .. } => {
if enough_groups {
self.state = State::Done;
} else if points.is_empty() {
self.state = State::Filling {
requests_left: fill_budget,
};
}
}
State::Filling { .. } => {
if enough_groups || points.is_empty() {
self.state = State::Done;
}
}
State::Done => {}
}
}
fn exclude_aggregated_points(&self, request: &mut ShardQueryRequest) {
let ids = self.aggregator.ids().clone();
if !ids.is_empty() {
merge_filter(
&mut request.filter,
Filter::new_must_not(Condition::HasId(ids.into())),
);
}
}
pub fn distill(self) -> Vec<Group> {
self.aggregator.distill()
}
}
#[cfg(test)]
mod tests {
use crate::common::types::ScoreType;
use crate::segment::data_types::groups::GroupId;
use crate::segment::payload_json;
use crate::segment::types::{WithPayloadInterface, WithVector};
use super::*;
const GROUPS: usize = 2;
const GROUP_SIZE: usize = 2;
fn base_query() -> ShardQueryRequest {
ShardQueryRequest {
prefetches: vec![],
query: None,
filter: None,
score_threshold: None,
limit: 0,
offset: 0,
params: None,
with_vector: WithVector::Bool(false),
with_payload: WithPayloadInterface::Bool(false),
}
}
fn driver(budget: RequestBudget) -> GroupByDriver {
GroupByDriver::new(
base_query(),
"g".parse().unwrap(),
GROUPS,
GROUP_SIZE,
Some(Order::LargeBetter),
budget,
)
}
fn point(id: u64, score: ScoreType, group: &str) -> ScoredPoint {
ScoredPoint {
id: id.into(),
version: 0,
score,
payload: Some(payload_json! { "g": group }),
vector: None,
shard_key: None,
order_value: None,
}
}
#[test]
fn single_request_budget_yields_one_shaped_request() {
let mut driver = driver(RequestBudget {
collect: 1,
fill: 0,
});
let request = driver.next_request().unwrap();
assert_eq!(request.limit, GROUPS * GROUP_SIZE);
assert_eq!(request.offset, 0);
assert!(request.filter.is_some());
assert_eq!(
request.with_payload,
WithPayloadInterface::Fields(vec!["g".parse().unwrap()])
);
driver.add_points(&[point(1, 1.0, "a")]);
assert!(driver.next_request().is_none());
let groups = driver.distill();
assert_eq!(groups.len(), 1);
assert_eq!(groups[0].key, GroupId::from("a"));
}
#[test]
fn stops_early_when_enough_groups_are_filled() {
let mut driver = driver(RequestBudget {
collect: 5,
fill: 5,
});
assert!(driver.next_request().is_some());
driver.add_points(&[
point(1, 4.0, "a"),
point(2, 3.0, "a"),
point(3, 2.0, "b"),
point(4, 1.0, "b"),
]);
assert!(driver.next_request().is_none());
let groups = driver.distill();
assert_eq!(groups.len(), GROUPS);
assert_eq!(groups[0].key, GroupId::from("a"));
assert_eq!(groups[1].key, GroupId::from("b"));
for group in &groups {
assert_eq!(group.hits.len(), GROUP_SIZE);
}
}
#[test]
fn moves_to_filling_and_finishes_on_empty_responses() {
let mut driver = driver(RequestBudget {
collect: 5,
fill: 5,
});
let collect_request = driver.next_request().unwrap();
driver.add_points(&[point(1, 2.0, "a"), point(2, 1.0, "b")]);
let second_request = driver.next_request().unwrap();
assert_ne!(collect_request.filter, second_request.filter);
driver.add_points(&[]);
assert!(driver.next_request().is_some());
driver.add_points(&[]);
assert!(driver.next_request().is_none());
let groups = driver.distill();
assert_eq!(groups.len(), 2);
}
#[test]
fn zero_groups_or_group_size_finishes_immediately() {
let budget = RequestBudget {
collect: 5,
fill: 5,
};
for (groups, group_size) in [(0, GROUP_SIZE), (GROUPS, 0)] {
let mut driver = GroupByDriver::new(
base_query(),
"g".parse().unwrap(),
groups,
group_size,
Some(Order::LargeBetter),
budget,
);
assert!(driver.next_request().is_none());
assert!(driver.distill().is_empty());
}
}
}