use std::collections::HashSet;
use gaoya::simhash::SimHashBits;
use crate::llmtrim::stages::dedup::make_simhasher;
use crate::llmtrim::stages::tools::{fnv1a, lex_words};
const KNEE_MIN_DEV: f64 = 0.05;
const SIMHASH_NEAR_DIST: u32 = 3;
pub fn optimal_keep(items: &[&str], min_k: usize, max_k: usize) -> usize {
let n = items.len();
let lo = min_k.min(n);
let hi = max_k.min(n);
if n <= lo {
return n;
}
if hi <= lo {
return lo;
}
let curve = unique_bigram_curve(items);
let knee = match knee_index(&curve) {
Some(i) => i + 1,
None if curve.last().copied().unwrap_or(0) == 0 => lo,
None => hi,
};
let clusters = distinct_clusters(items).max(lo);
knee.clamp(lo, hi).min(clusters)
}
fn unique_bigram_curve(items: &[&str]) -> Vec<usize> {
let mut seen: HashSet<u64> = HashSet::new();
let mut curve = Vec::with_capacity(items.len());
for item in items {
let words = lex_words(item);
match words.as_slice() {
[] => {}
[single] => {
seen.insert(token_hash(single, ""));
}
_ => {
for pair in words.windows(2) {
seen.insert(token_hash(&pair[0], &pair[1]));
}
}
}
curve.push(seen.len());
}
curve
}
fn knee_index(curve: &[usize]) -> Option<usize> {
let n = curve.len();
if n < 3 {
return None;
}
let total = *curve.last().unwrap() as f64;
if total <= 0.0 {
return None;
}
let x_span = (n - 1) as f64;
let mut best = (0usize, 0.0f64);
for (i, &c) in curve.iter().enumerate() {
let x = i as f64 / x_span;
let y = c as f64 / total;
let deviation = y - x;
if deviation > best.1 {
best = (i, deviation);
}
}
(best.1 > KNEE_MIN_DEV).then_some(best.0)
}
fn distinct_clusters(items: &[&str]) -> usize {
let hasher = make_simhasher();
let mut reps: Vec<u64> = Vec::new();
for item in items {
let words = lex_words(item);
let sig = if words.is_empty() {
0
} else {
hasher.create_signature(words.iter())
};
if reps
.iter()
.any(|&r| r.hamming_distance(&sig) <= SIMHASH_NEAR_DIST as usize)
{
continue;
}
reps.push(sig);
}
reps.len()
}
fn token_hash(a: &str, b: &str) -> u64 {
fnv1a(a.bytes().chain(std::iter::once(0x1f)).chain(b.bytes()))
}
#[cfg(test)]
mod tests {
use super::*;
fn diverse(n: usize) -> Vec<String> {
(0..n)
.map(|i| {
(0..6)
.map(|j| format!("w{i}x{j}"))
.collect::<Vec<_>>()
.join(" ")
})
.collect()
}
fn as_refs(v: &[String]) -> Vec<&str> {
v.iter().map(String::as_str).collect()
}
#[test]
fn returns_all_when_under_floor() {
let items = ["a", "b"];
assert_eq!(optimal_keep(&items, 5, 10), 2, "n below min_k keeps all");
assert_eq!(optimal_keep(&[], 1, 10), 0, "empty keeps none");
}
#[test]
fn diverse_content_keeps_up_to_budget() {
let v = diverse(20);
let k = optimal_keep(&as_refs(&v), 3, 12);
assert_eq!(k, 12, "all-diverse curve has no knee → keep the max budget");
}
#[test]
fn near_duplicate_spam_is_capped_by_clusters() {
let v: Vec<String> =
std::iter::repeat_n("WARN cache miss for user session".to_string(), 20).collect();
let k = optimal_keep(&as_refs(&v), 2, 15);
assert!(
k <= 3,
"20 near-identical lines collapse to ~1 cluster, got {k}"
);
}
#[test]
fn early_saturation_keeps_fewer_than_diverse() {
let mut v = diverse(4);
v.extend(std::iter::repeat_n(
"retry pending retry pending retry".to_string(),
16,
));
let saturating = optimal_keep(&as_refs(&v), 2, 18);
let all_diverse = optimal_keep(&as_refs(&diverse(20)), 2, 18);
assert!(
saturating < all_diverse,
"saturating set ({saturating}) keeps fewer than diverse ({all_diverse})"
);
}
#[test]
fn never_exceeds_item_count() {
let v = diverse(6);
assert!(optimal_keep(&as_refs(&v), 2, 100) <= 6);
}
#[test]
fn knee_on_concave_curve_is_near_the_elbow() {
let curve = [0usize, 5, 9, 12, 13, 13, 13, 13];
let knee = knee_index(&curve).expect("a concave curve has a knee");
assert!(
(2..=4).contains(&knee),
"knee at the elbow region, got {knee}"
);
}
#[test]
fn linear_curve_has_no_knee() {
let curve = [0usize, 2, 4, 6, 8, 10];
assert_eq!(knee_index(&curve), None, "a straight line has no knee");
}
}