From b7d8cda047df47ee13cf662ccce1eba8a434cadc Mon Sep 17 00:00:00 2001 From: dawn <90008@klbr.net> Date: Sat, 27 Jun 2026 20:39:02 +0300 Subject: [PATCH] feat: implement principled scoring and adaptive neighbor expansion --- klbr-core/src/evidence.rs | 135 +++++++++++++++++++++++++++++++------- klbr-core/src/pipeline.rs | 3 + 2 files changed, 113 insertions(+), 25 deletions(-) diff --git a/klbr-core/src/evidence.rs b/klbr-core/src/evidence.rs index 7cdff45..1622f61 100644 --- a/klbr-core/src/evidence.rs +++ b/klbr-core/src/evidence.rs @@ -184,6 +184,47 @@ pub struct EvidenceOmittedRef { pub reason: String, } +fn query_is_wh_like(query: &str) -> bool { + let lower = query.to_lowercase(); + ["what", "who", "where", "why", "how", "when", "which"] + .iter() + .any(|word| lower.starts_with(word) || lower.contains(&format!(" {word}"))) +} + +fn looks_incomplete(body: &str) -> bool { + let trimmed = body.trim(); + trimmed.len() < 30 + || trimmed.ends_with(':') + || trimmed.ends_with(',') + || ["he", "she", "they", "it", "that", "this"].iter().any(|pronoun| trimmed.to_lowercase().contains(pronoun)) +} + +fn neighbor_gain(neighbor_body: &str, anchor_body: &str, query: &str) -> f32 { + let lower_neighbor = neighbor_body.to_lowercase(); + let lower_query = query.to_lowercase(); + + let mut fts_score = 0.0; + for word in lower_query.split_whitespace() { + if word.len() >= 4 && lower_neighbor.contains(word) { + fts_score += 1.0; + } + } + + let mut entity_score = 0.0; + for word in anchor_body.split_whitespace() { + if let Some(first_char) = word.chars().next() { + if first_char.is_uppercase() && word.len() >= 3 { + let clean_word = word.trim_matches(|c: char| !c.is_alphabetic()).to_lowercase(); + if !clean_word.is_empty() && lower_neighbor.contains(&clean_word) { + entity_score += 1.5; + } + } + } + } + + fts_score + entity_score +} + #[derive(Clone)] pub struct EvidencePlanner { memory: MemoryStore, @@ -196,6 +237,7 @@ impl EvidencePlanner { pub fn build_packets( &self, + query: &str, candidates: &[EvidenceAtom], limit: usize, packet_token_budget: usize, @@ -206,7 +248,7 @@ impl EvidencePlanner { for (index, candidate) in candidates.iter().enumerate() { let expanded = - self.expand_candidate_packet(candidate, index + 1, packet_token_budget)?; + self.expand_candidate_packet(query, candidate, index + 1, packet_token_budget)?; let packet_id = expanded.packet.packet_id.clone(); omissions.extend(expanded.omissions); if let Some(position) = packet_positions.get(&packet_id).copied() { @@ -227,6 +269,7 @@ impl EvidencePlanner { fn expand_candidate_packet( &self, + query: &str, candidate: &EvidenceAtom, rank: usize, packet_token_budget: usize, @@ -249,8 +292,26 @@ impl EvidencePlanner { let expansion_refs = match packet_kind { EvidencePacketKind::TurnWindow => { - self.memory - .turn_window_ref_ids(&candidate.ref_id, 2, 2, 18)? + let tight_refs = self.memory + .turn_window_ref_ids(&candidate.ref_id, 1, 1, 10)?; + let mut final_refs = tight_refs; + if query_is_wh_like(query) || looks_incomplete(&candidate.body) { + let wider_refs = self.memory + .turn_window_ref_ids(&candidate.ref_id, 2, 2, 18)?; + let extra_refs: Vec = wider_refs.into_iter() + .filter(|r| !final_refs.contains(r)) + .collect(); + if !extra_refs.is_empty() { + let extra_entries = self.memory + .promptable_entries_for_refs(&extra_refs, "turn_window")?; + for entry in extra_entries { + if neighbor_gain(&entry.body, &candidate.body, query) >= 1.0 { + final_refs.push(entry.ref_id); + } + } + } + } + final_refs } EvidencePacketKind::ExactRef | EvidencePacketKind::EpisodeBridge @@ -511,20 +572,54 @@ pub(crate) fn merge_packet_signals(existing: &mut EvidencePacket, incoming: Evid } fn packet_fusion_score(packet: &EvidencePacket) -> f32 { - let exact_bonus = if packet.packet_kind == EvidencePacketKind::ExactRef { - 100.0 + let rrf_k = 20.0; + + let exact_rank = packet.signals.per_source.iter() + .filter(|sig| sig.source == "exact") + .map(|sig| sig.rank) + .min(); + let fts_rank = packet.signals.per_source.iter() + .filter(|sig| sig.source == "fts") + .map(|sig| sig.rank) + .min(); + let dense_rank = packet.signals.per_source.iter() + .filter(|sig| sig.source == "dense") + .map(|sig| sig.rank) + .min(); + let graph_rank = packet.signals.per_source.iter() + .filter(|sig| sig.source == "graph" || sig.source == "graph_support") + .map(|sig| sig.rank) + .min(); + + let mut rrf_score = 0.0; + if let Some(r) = exact_rank { + rrf_score += 3.0 / (rrf_k + r as f32); + } + if let Some(r) = dense_rank { + rrf_score += 1.2 / (rrf_k + r as f32); + } + if let Some(r) = fts_rank { + rrf_score += 1.0 / (rrf_k + r as f32); + } + if let Some(r) = graph_rank { + rrf_score += 0.6 / (rrf_k + r as f32); + } + + let source_support_count = packet.signals.sources.len() as f32; + let support_bonus = 0.25 * (1.0 + source_support_count).ln(); + + let explicit_bonus = if packet.packet_kind == EvidencePacketKind::ExactRef { + 5.0 } else { 0.0 }; - let signal_score = packet - .signals - .per_source - .iter() - .map(|signal| source_weight(&signal.source) / signal.rank.max(1) as f32) - .sum::(); - let agreement_bonus = packet.signals.sources.len().saturating_sub(1) as f32 * 0.25; - let support_bonus = packet.refs.len().saturating_sub(1) as f32 * 0.05; - exact_bonus + signal_score + agreement_bonus + support_bonus + + let has_neighbor = packet.packet_kind == EvidencePacketKind::TurnWindow && packet.refs.len() > 1; + let neighbor_bonus = if has_neighbor { 1.0 } else { 0.0 }; + + let token_penalty = 0.1 * (1.0 + packet.estimated_tokens as f32).ln(); + + rrf_score + support_bonus + explicit_bonus + neighbor_bonus - token_penalty } fn packet_ref_overlap(left: &EvidencePacket, right: &EvidencePacket) -> f32 { @@ -538,17 +633,7 @@ fn packet_ref_overlap(left: &EvidencePacket, right: &EvidencePacket) -> f32 { intersection as f32 / union as f32 } -fn source_weight(source: &str) -> f32 { - match source { - "exact" => 10.0, - "fts" => 4.0, - "dense" => 4.0, - "graph" => 2.0, - "turn_window" | "exact_support" | "episode_support" | "graph_support" => 0.5, - "complete_stored" => 1.0, - _ => 1.0, - } -} + fn source_priority(source: &str) -> usize { match source { diff --git a/klbr-core/src/pipeline.rs b/klbr-core/src/pipeline.rs index 7845f3c..526a02b 100644 --- a/klbr-core/src/pipeline.rs +++ b/klbr-core/src/pipeline.rs @@ -344,6 +344,7 @@ impl MemoryPipeline { let candidates = rank_seed_candidates(candidates, budget.top_k.saturating_mul(4).max(budget.top_k)); let mut packet_plan = self.evidence_planner().build_packets( + &query.text, &candidates, budget.top_k, budget.max_tokens.min(900), @@ -383,6 +384,7 @@ impl MemoryPipeline { fallback_packets = self .evidence_planner() .build_packets( + &query.text, &retrieved.candidates, budget.top_k, budget.max_tokens.min(900), @@ -447,6 +449,7 @@ impl MemoryPipeline { .map(|entry| self.ref_search_entry_to_atom(entry)) .collect::>>()?; let mut packet_plan = self.evidence_planner().build_packets( + &query.text, &candidates, budget.top_k, budget.max_tokens.min(900), -- 2.51.2