use std::collections::{HashMap, HashSet};
use anyhow::Result;
use once_cell::sync::Lazy;
use regex::Regex;
use serde_json::Value;
use crate::llmtrim::gate::{GateKind, PlanEntry, Scope, Transform};
use crate::llmtrim::ir::Request;
use crate::llmtrim::provider::Provider;
use crate::llmtrim::select::{self, Item, Weights};
use crate::llmtrim::stages::tools::lex_words;
const SAMPLE_NOTE: &str = include_str!("../prompts/jsoncrush_note.txt");
static ERROR_KEYWORD: Lazy<Regex> = Lazy::new(|| {
Regex::new(r"(?i)\b(error|fail(?:ed|ure)?|fatal|panic|exception|denied|invalid|timeout)\b")
.unwrap()
});
pub struct JsonCrushStage {
pub max_rows: usize,
}
impl Transform for JsonCrushStage {
fn name(&self) -> &str {
"json-crush"
}
fn gate_kind(&self) -> GateKind {
GateKind::InputTokens
}
fn scope(&self) -> Scope {
Scope::Content
}
fn apply(
&self,
req: &mut Request,
provider: &dyn Provider,
_plan: &mut Vec<PlanEntry>,
) -> Result<()> {
let pointers = crate::llmtrim::cache_zone::compressible_pointers(req, provider);
let query: HashSet<String> = pointers
.iter()
.filter_map(|p| req.get_str(p))
.filter(|t| t.lines().count() < 4 && t.len() < 600)
.flat_map(lex_words)
.collect();
let mut sampled_any = false;
let mut bare_counts: Vec<usize> = Vec::new();
for ptr in &pointers {
let Some(s) = req.get_str(ptr).map(str::to_string) else {
continue;
};
let Ok(mut value) = serde_json::from_str::<Value>(&s) else {
continue;
};
if let Some(bare_n) = crush_value(&mut value, self.max_rows, &query) {
req.set(ptr, Value::String(value.to_string()));
sampled_any = true;
bare_counts.extend(bare_n);
}
}
if sampled_any {
provider.add_system_instruction(req, &sample_note(&bare_counts));
}
Ok(())
}
}
fn sample_note(bare_counts: &[usize]) -> String {
if bare_counts.is_empty() {
return SAMPLE_NOTE.to_string();
}
let counts = bare_counts
.iter()
.map(usize::to_string)
.collect::<Vec<_>>()
.join(", ");
let plural = if bare_counts.len() == 1 {
"array's original row count was"
} else {
"arrays' original row counts were"
};
format!("{} The sampled {plural}: {counts}.", SAMPLE_NOTE.trim_end())
}
fn crush_value(
value: &mut Value,
max_rows: usize,
query: &HashSet<String>,
) -> Option<Option<usize>> {
if let Some((n, rows)) = crush_array(value, max_rows, query) {
*value = Value::Array(rows);
return Some(Some(n));
}
if let Value::Object(map) = value {
let mut sampled: Vec<(String, usize)> = Vec::new();
for (key, field) in map.iter_mut() {
if let Some((n, rows)) = crush_array(field, max_rows, query) {
*field = Value::Array(rows);
sampled.push((key.clone(), n));
}
}
if sampled.is_empty() {
return None;
}
for (key, n) in sampled {
map.insert(format!("_sampled_from_{key}"), Value::from(n));
}
return Some(None);
}
None
}
fn crush_array(v: &Value, max_rows: usize, query: &HashSet<String>) -> Option<(usize, Vec<Value>)> {
let arr = v.as_array()?;
let n = arr.len();
if n <= max_rows || !arr.iter().all(Value::is_object) {
return None;
}
let serialized: Vec<String> = arr.iter().map(Value::to_string).collect();
let k_first = ((max_rows as f64 * 0.6).round() as usize).clamp(1, n);
let k_last = ((max_rows as f64 * 0.2).round() as usize).clamp(1, (n - k_first).max(1));
let mut keep = vec![false; n];
for slot in keep.iter_mut().take(k_first) {
*slot = true;
}
for slot in keep.iter_mut().skip(n - k_last) {
*slot = true;
}
let mut count = keep.iter().filter(|&&x| x).count();
for &i in outlier_rows(arr, &serialized).iter() {
if count >= max_rows {
break;
}
if !keep[i] {
keep[i] = true;
count += 1;
}
}
fill_diverse(arr, &serialized, &mut keep, max_rows, query);
if keep.iter().filter(|&&k| k).count() >= n {
return None;
}
let rows = arr
.iter()
.zip(&keep)
.filter(|&(_, &k)| k)
.map(|(row, _)| row.clone())
.collect();
Some((n, rows))
}
fn fill_diverse(
arr: &[Value],
serialized: &[String],
keep: &mut [bool],
max_rows: usize,
query: &HashSet<String>,
) {
let used = keep.iter().filter(|&&k| k).count();
let remaining = max_rows.saturating_sub(used);
if remaining == 0 {
return;
}
let candidates: Vec<usize> = (0..arr.len()).filter(|&i| !keep[i]).collect();
let items: Vec<Item> = candidates
.iter()
.map(|&i| {
let rel = query_overlap(&serialized[i], query);
Item::from_text(&row_value_text(&arr[i]), 1, rel)
})
.collect();
for local in select::select(&items, remaining, &Weights::default()) {
keep[candidates[local]] = true;
}
}
fn row_value_text(row: &Value) -> String {
let mut out = String::new();
collect_scalar_values(row, &mut out);
out
}
fn collect_scalar_values(v: &Value, out: &mut String) {
match v {
Value::String(s) => {
out.push_str(s);
out.push(' ');
}
Value::Number(_) | Value::Bool(_) => {
out.push_str(&v.to_string());
out.push(' ');
}
Value::Array(a) => {
for e in a {
collect_scalar_values(e, out);
}
}
Value::Object(m) => {
for val in m.values() {
collect_scalar_values(val, out);
}
}
Value::Null => {}
}
}
fn outlier_rows(arr: &[Value], serialized: &[String]) -> Vec<usize> {
let n = arr.len();
let rare_at = (n / 20).max(1);
let cat_cap = (n / 10).clamp(2, 24);
let mut freq: HashMap<&str, HashMap<String, usize>> = HashMap::new();
for row in arr {
if let Some(obj) = row.as_object() {
for (key, val) in obj {
if is_scalar(val) {
*freq
.entry(key.as_str())
.or_default()
.entry(val.to_string())
.or_default() += 1;
}
}
}
}
let mut out = Vec::new();
for (i, row) in arr.iter().enumerate() {
if ERROR_KEYWORD.is_match(&serialized[i]) {
out.push(i);
continue;
}
let Some(obj) = row.as_object() else { continue };
for (key, val) in obj {
if !is_scalar(val) {
continue;
}
if let Some(counts) = freq.get(key.as_str()) {
let distinct = counts.len();
if (2..=cat_cap).contains(&distinct)
&& counts.get(&val.to_string()).copied().unwrap_or(0) <= rare_at
{
out.push(i);
break;
}
}
}
}
out
}
fn is_scalar(v: &Value) -> bool {
!v.is_array() && !v.is_object()
}
fn query_overlap(row: &str, query: &HashSet<String>) -> f64 {
if query.is_empty() {
return 0.0;
}
lex_words(row)
.into_iter()
.filter(|w| query.contains(w))
.count() as f64
}
#[cfg(test)]
mod tests {
use super::*;
use crate::llmtrim::ir::ProviderKind;
use crate::llmtrim::pipeline;
use crate::llmtrim::provider::OpenAiProvider;
use crate::llmtrim::tokenizer::counter_for;
use serde_json::json;
fn records(n: usize) -> Value {
let mut a = Vec::new();
for i in 0..n {
let status = if i == 7 || i == 900 { "error" } else { "ok" };
a.push(json!({"id": i, "status": status, "msg": format!("request {i} handled")}));
}
Value::Array(a)
}
#[test]
fn samples_big_array_and_keeps_outliers() {
let arr = records(1000);
let q = HashSet::new();
let (n, rows) = crush_array(&arr, 50, &q).expect("over-cap array is sampled");
assert_eq!(n, 1000, "original row count reported");
assert!(rows.len() <= 50, "down to the budget, got {}", rows.len());
let errors = rows.iter().filter(|r| r["status"] == "error").count();
assert_eq!(errors, 2, "rare error rows are kept as outliers");
assert_eq!(rows.first().unwrap()["id"], 0);
assert_eq!(rows.last().unwrap()["id"], 999);
}
#[test]
fn outliers_are_capped_to_budget() {
let arr = Value::Array(
(0..1000)
.map(|i| json!({"id": i, "status": "error", "msg": format!("fail {i}")}))
.collect(),
);
let q = HashSet::new();
let (_, rows) = crush_array(&arr, 50, &q).expect("over-cap array is sampled");
assert!(
rows.len() <= 50,
"outliers bounded by budget, got {}",
rows.len()
);
}
#[test]
fn all_error_array_drops_rows_instead_of_keeping_everything() {
let n = 1000;
let arr = Value::Array(
(0..n)
.map(|i| json!({"id": i, "status": "error"}))
.collect(),
);
let q = HashSet::new();
let (_, rows) = crush_array(&arr, 50, &q).expect("over-cap array is sampled");
assert!(rows.len() < n, "rows actually dropped, not all kept");
assert!(rows.len() <= 50, "within budget");
}
#[test]
fn small_or_scalar_arrays_are_left_alone() {
let q = HashSet::new();
assert!(
crush_array(&records(20), 50, &q).is_none(),
"below cap → serialize's job"
);
assert!(
crush_array(&json!([1, 2, 3, 4, 5]), 2, &q).is_none(),
"scalar array → not a record array"
);
}
#[test]
fn stage_reduces_tokens_on_a_huge_array() {
let content = serde_json::to_string(&records(1000)).unwrap();
let body = json!({"model": "gpt-4o", "messages": [{"role": "user", "content": content}], "max_tokens": 100});
let mut req = Request::from_value(ProviderKind::OpenAi, body);
let counter = counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap();
let stages: Vec<Box<dyn Transform>> = vec![Box::new(JsonCrushStage { max_rows: 50 })];
let out = pipeline::run(&mut req, &OpenAiProvider, counter.as_ref(), &stages);
assert!(out.stages[0].applied, "1000-row array sampled");
assert!(out.input_tokens_after < out.input_tokens_before);
let encoded = req.get_str("/messages/1/content").unwrap();
assert!(
encoded.contains("\"error\""),
"error rows survive the sample"
);
}
#[test]
fn nested_array_field_is_sampled_in_place() {
let wrapper = json!({"results": records(1000), "total": 1000});
let mut v = wrapper;
assert_eq!(
crush_value(&mut v, 50, &HashSet::new()),
Some(None),
"nested array sampled (no bare count surfaced in the note)"
);
assert_eq!(v["total"], 1000, "sibling fields preserved");
assert!(
v["results"].as_array().unwrap().len() <= 50,
"results sampled"
);
}
#[test]
fn nested_array_gets_sampled_from_sibling_with_count() {
let mut v = json!({"results": records(1000)});
assert_eq!(
crush_value(&mut v, 50, &HashSet::new()),
Some(None),
"nested array sampled"
);
assert_eq!(
v["_sampled_from_results"], 1000,
"sibling carries the original row count"
);
assert!(
v["results"].as_array().unwrap().len() <= 50,
"results sampled, still an array of objects"
);
assert!(v["results"].is_array(), "field stays an array");
}
#[test]
fn bare_top_level_array_surfaces_count_in_note() {
let mut v = records(1000);
assert_eq!(
crush_value(&mut v, 50, &HashSet::new()),
Some(Some(1000)),
"bare array reports its original row count"
);
assert!(v.is_array(), "stays a bare array");
assert!(v.as_array().unwrap().len() <= 50, "sampled");
let note = sample_note(&[1000]);
assert!(note.contains("1000"), "note surfaces N: {note}");
}
#[test]
fn sample_note_is_static_without_bare_counts() {
assert_eq!(
sample_note(&[]),
SAMPLE_NOTE,
"no bare counts → unchanged static note"
);
}
#[test]
fn sample_note_uses_plural_wording_for_multiple_bare_counts() {
let note = sample_note(&[100, 200]);
assert!(
note.contains("arrays' original row counts were"),
"plural wording: {note}"
);
assert!(
note.contains("100, 200"),
"both counts comma-joined: {note}"
);
}
#[test]
fn multiple_nested_array_fields_each_get_their_own_sibling() {
let mut v = json!({"a": records(1000), "b": records(800)});
assert_eq!(
crush_value(&mut v, 50, &HashSet::new()),
Some(None),
"only nested fields sampled"
);
assert_eq!(v["_sampled_from_a"], 1000, "field a's original count");
assert_eq!(v["_sampled_from_b"], 800, "field b's original count");
assert!(
v["a"].as_array().unwrap().len() <= 50,
"a sampled, still array"
);
assert!(
v["b"].as_array().unwrap().len() <= 50,
"b sampled, still array"
);
}
#[test]
fn mixed_bare_and_nested_arrays_accumulate_independently() {
let mut bare = records(1000);
let mut nested = json!({"rows": records(900)});
let mut bare_counts: Vec<usize> = Vec::new();
bare_counts.extend(crush_value(&mut bare, 50, &HashSet::new()).flatten());
bare_counts.extend(crush_value(&mut nested, 50, &HashSet::new()).flatten());
assert_eq!(
bare_counts,
vec![1000],
"only the bare array contributes a note count"
);
assert!(bare.is_array(), "bare stays a bare array");
assert_eq!(
nested["_sampled_from_rows"], 900,
"nested array gets its sibling"
);
}
#[test]
fn diverse_fill_prefers_distinct_rows_over_near_duplicate_spam() {
let mut a: Vec<Value> = Vec::new();
for _ in 0..120 {
a.push(json!({"kind": "x", "msg": "routine heartbeat ping ok steady nominal"}));
}
let distinct = [
"disk volume remount latency spike detected",
"auth token rotation completed for tenant",
"cache warm reload finished across shards",
"queue backlog drained after worker scale",
"tls handshake renegotiated upstream peer",
];
let pos: Vec<usize> = (0..distinct.len()).map(|k| 40 + k * 3).collect();
for (k, &p) in pos.iter().enumerate() {
a[p] = json!({"kind": "x", "msg": distinct[k]});
}
let arr = Value::Array(a);
let (_, rows) = crush_array(&arr, 30, &HashSet::new()).expect("over-cap array is sampled");
let msgs: HashSet<&str> = rows.iter().filter_map(|r| r["msg"].as_str()).collect();
let distinct_kept = distinct.iter().filter(|d| msgs.contains(**d)).count();
assert!(
distinct_kept >= 3,
"diverse fill surfaces the distinct rows (kept {distinct_kept}/5): {msgs:?}"
);
assert!(rows.len() <= 30, "within budget, got {}", rows.len());
}
#[test]
fn query_bias_survives_diverse_fill() {
let mut a: Vec<Value> = Vec::new();
for i in 0..400 {
a.push(json!({"kind": "x", "msg": format!("routine event number {i}")}));
}
a[200] = json!({"kind": "x", "msg": "kubernetes pod eviction quota exceeded"});
let arr = Value::Array(a);
let query: HashSet<String> = lex_words("kubernetes pod eviction").into_iter().collect();
let (_, rows) = crush_array(&arr, 30, &query).expect("over-cap array is sampled");
let kept_needle = rows
.iter()
.any(|r| r["msg"].as_str() == Some("kubernetes pod eviction quota exceeded"));
assert!(
kept_needle,
"the query-matching row is kept (relevance term)"
);
}
#[test]
fn row_value_text_uses_values_not_keys() {
let a = row_value_text(&json!({"city": "Paris", "code": 75}));
let b = row_value_text(&json!({"city": "Tokyo", "code": 13}));
assert!(
a.contains("Paris") && a.contains("75"),
"values present: {a:?}"
);
assert!(!a.contains("city"), "keys excluded: {a:?}");
assert_ne!(a, b, "different values → different feature text");
}
}