use std::collections::HashMap;
use kevy_text::cold::{decode_fwd, posting_df, score_cold, score_cold_phrase};
use kevy_text::{CorpusStats, sorted_order};
use super::TextColdDir;
#[derive(Debug)]
pub struct ColdPageQuery<'a> {
pub bare: Vec<Vec<u8>>,
pub phrases: Vec<Vec<Vec<u8>>>,
pub stats: &'a CorpusStats,
pub filter: &'a [kevy_text::Filter<'a>],
pub sort: Option<&'a kevy_text::Sort<'a>>,
pub distinct: Option<&'a kevy_text::Distinct<'a>>,
pub facets: &'a [kevy_text::Facet<'a>],
pub fetch: usize,
}
#[derive(Debug)]
pub struct ColdHit {
pub key: Vec<u8>,
pub score: f64,
pub okey: Option<Vec<u8>>,
}
#[derive(Debug)]
pub struct ColdPage {
pub hits: Vec<ColdHit>,
pub values: HashMap<Vec<u8>, Vec<Option<Vec<u8>>>>,
pub facets: Vec<Vec<kevy_text::Bucket>>,
}
impl TextColdDir {
pub fn cold_stats(&self, tokens: &[Vec<u8>]) -> (u64, u64, Vec<(Vec<u8>, u32)>) {
let n_docs: u64 = self.segs.iter().map(|c| c.n_docs).sum();
let total_len: u64 = self.segs.iter().map(|c| c.total_len).sum();
let df = tokens
.iter()
.map(|t| {
let frozen: u32 = self
.segs
.iter()
.filter_map(|c| c.seg.get(t).ok().flatten())
.filter_map(|p| posting_df(&p))
.sum();
let dead = self.df_dead.get(t).copied().unwrap_or(0);
(t.clone(), frozen.saturating_sub(dead))
})
.collect();
(n_docs, total_len, df)
}
pub fn cold_page(&self, q: &ColdPageQuery) -> ColdPage {
let acc = self.accumulate(q);
let need_values = !q.filter.is_empty()
|| q.sort.is_some()
|| q.distinct.is_some()
|| !q.facets.is_empty();
let mut values: HashMap<Vec<u8>, Vec<Option<Vec<u8>>>> = HashMap::new();
let mut cands: Vec<ColdHit> = Vec::new();
for (key, score) in acc {
let vals = if need_values {
let Some(v) = self.frozen_values(&key) else { continue };
if !passes(&v, q.filter) {
continue;
}
Some(v)
} else {
None
};
let okey = q.sort.and_then(|s| vals.as_ref()?.get(s.field)?.as_deref().and_then(s.key));
if let Some(v) = vals {
values.insert(key.clone(), v);
}
cands.push(ColdHit { key, score, okey });
}
let facets = self.count_facets(q, &cands, &values);
order_page(&mut cands, q.sort.is_some(), q.sort.is_some_and(|s| s.desc));
if let Some(d) = q.distinct {
collapse(&mut cands, d, &values);
}
cands.truncate(q.fetch);
values.retain(|k, _| cands.iter().any(|c| &c.key == k));
ColdPage { hits: cands, values, facets }
}
fn accumulate(&self, q: &ColdPageQuery) -> HashMap<Vec<u8>, f64> {
let mut acc = HashMap::new();
for cs in &self.segs {
let dead = |k: &[u8]| self.tombs.get(k).is_some_and(|s| s.contains(&cs.seq));
for t in &q.bare {
if let Ok(Some(payload)) = cs.seg.get(t) {
let _ = score_cold(&payload, t, q.stats, &dead, &mut acc);
}
}
for phrase in &q.phrases {
let payloads: Option<Vec<Vec<u8>>> =
phrase.iter().map(|t| cs.seg.get(t).ok().flatten()).collect();
if let Some(payloads) = payloads {
let _ = score_cold_phrase(&payloads, phrase, q.stats, &dead, &mut acc);
}
}
}
acc
}
fn frozen_values(&self, key: &[u8]) -> Option<Vec<Option<Vec<u8>>>> {
let mut fwd_key = vec![0u8];
fwd_key.extend_from_slice(key);
for cs in &self.segs {
if self.tombs.get(key).is_some_and(|s| s.contains(&cs.seq)) {
continue;
}
if let Ok(Some(payload)) = cs.seg.get(&fwd_key) {
return decode_fwd(&payload).map(|r| r.values);
}
}
None
}
fn count_facets(
&self,
q: &ColdPageQuery,
cands: &[ColdHit],
values: &HashMap<Vec<u8>, Vec<Option<Vec<u8>>>>,
) -> Vec<Vec<kevy_text::Bucket>> {
q.facets
.iter()
.map(|f| {
let mut counts: HashMap<Vec<u8>, (Vec<u8>, u64)> = HashMap::new();
for c in cands {
let Some(raw) =
values.get(&c.key).and_then(|v| v.get(f.field)).and_then(Option::as_deref)
else {
continue;
};
let Some(k) = (f.key)(raw) else { continue };
counts.entry(k).or_insert_with(|| (raw.to_vec(), 0)).1 += 1;
}
let mut out: Vec<kevy_text::Bucket> =
counts.into_iter().map(|(k, (label, n))| (k, label, n)).collect();
out.sort_by(|a, b| b.2.cmp(&a.2).then_with(|| a.1.cmp(&b.1)));
out
})
.collect()
}
}
fn order_page(cands: &mut [ColdHit], sorted: bool, desc: bool) {
if sorted {
cands.sort_by(|a, b| {
sorted_order((a.okey.as_deref(), &a.key), (b.okey.as_deref(), &b.key), desc)
});
} else {
cands.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.key.cmp(&b.key))
});
}
}
fn collapse(
cands: &mut Vec<ColdHit>,
d: &kevy_text::Distinct,
values: &HashMap<Vec<u8>, Vec<Option<Vec<u8>>>>,
) {
let mut seen: std::collections::HashSet<Vec<u8>> = std::collections::HashSet::new();
cands.retain(|c| {
let identity = values
.get(&c.key)
.and_then(|v| v.get(d.field))
.and_then(Option::as_deref)
.and_then(d.key);
match identity {
None => true,
Some(id) => seen.insert(id),
}
});
}
fn passes(values: &[Option<Vec<u8>>], filter: &[kevy_text::Filter]) -> bool {
filter
.iter()
.all(|f| values.get(f.field).and_then(Option::as_deref).is_some_and(|v| (f.test)(v)))
}