use super::*;
const TFIDF_SCHEMA_VERSION: u32 = 1;
#[derive(Debug, Clone, Serialize, Deserialize)]
struct TfIdfPersistedState {
#[serde(default)]
schema_version: u32,
vocab: Vec<String>,
idf: Vec<f32>,
dimension: usize,
pdg_nodes: usize,
pdg_edges: usize,
#[serde(default)]
pdg_fingerprint: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TfIdfEmbedder {
pub(crate) vocab: Vec<String>,
pub(crate) idf: Vec<f32>,
pub(crate) dimension: usize,
pub(crate) pdg_nodes: usize,
pub(crate) pdg_edges: usize,
pub(crate) pdg_fingerprint: String,
}
impl TfIdfEmbedder {
#[cfg_attr(not(test), allow(dead_code))]
pub fn build(documents: &[(String, String)]) -> Self {
let tokenized: Vec<(String, Vec<String>)> = documents
.iter()
.map(|(id, content)| (id.clone(), tokenize_code(content)))
.collect();
Self::build_from_tokens(&tokenized)
}
pub fn build_from_tokens(documents: &[(String, Vec<String>)]) -> Self {
const TARGET_DIM: usize = crate::search::search::DEFAULT_EMBEDDING_DIMENSION;
let n = documents.len();
if n == 0 {
return Self {
vocab: Vec::new(),
idf: Vec::new(),
dimension: TARGET_DIM,
pdg_nodes: 0,
pdg_edges: 0,
pdg_fingerprint: String::new(),
};
}
let mut df: std::collections::HashMap<String, usize> = std::collections::HashMap::new();
let mut seen: std::collections::HashSet<&str> = std::collections::HashSet::new();
for (_, tokens) in documents {
seen.clear();
for tok in tokens {
if seen.insert(tok.as_str()) {
*df.entry(tok.to_string()).or_insert(0) += 1;
}
}
}
let n_f = n as f32;
let min_df: usize = if n < 50 { 1 } else { (n / 1000).max(3) };
let max_df: usize = if n < 50 { n } else { (n / 4).max(min_df + 1) };
let mut idf_scores: Vec<(String, f32)> = df
.into_iter()
.filter(|(_, df_count)| *df_count >= min_df && *df_count <= max_df)
.map(|(tok, df_count)| {
let idf = (n_f / df_count as f32).ln();
(tok, idf)
})
.collect();
info!(
vocab_candidates = idf_scores.len(),
min_df,
max_df,
n_docs = n,
"TF-IDF vocabulary candidates (moderate-IDF filter)"
);
idf_scores.sort_by(|a, b| {
a.1.partial_cmp(&b.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
let final_scores: Vec<(String, f32)> = if idf_scores.len() <= TARGET_DIM {
idf_scores
} else {
let total = idf_scores.len();
let stride = total as f64 / TARGET_DIM as f64;
(0..TARGET_DIM)
.map(|i| {
let idx = ((i as f64 * stride) as usize).min(total - 1);
idf_scores[idx].clone()
})
.collect()
};
let idf_scores = final_scores;
let vocab: Vec<String> = idf_scores.iter().map(|(t, _)| t.clone()).collect();
let idf: Vec<f32> = idf_scores.iter().map(|(_, s)| *s).collect();
Self {
vocab,
idf,
dimension: TARGET_DIM,
pdg_nodes: 0,
pdg_edges: 0,
pdg_fingerprint: String::new(),
}
}
pub fn embed(&self, text: &str) -> Vec<f32> {
let mut vec = vec![0.0f32; self.dimension];
if self.vocab.is_empty() {
return vec;
}
let tokens = tokenize_code(text);
let total = tokens.len() as f32;
if total == 0.0 {
return vec;
}
let mut tf_map: std::collections::HashMap<&str, f32> = std::collections::HashMap::new();
for tok in &tokens {
*tf_map.entry(tok.as_str()).or_insert(0.0) += 1.0;
}
for (slot, (word, idf_val)) in vec.iter_mut().zip(self.vocab.iter().zip(self.idf.iter())) {
if let Some(&count) = tf_map.get(word.as_str()) {
*slot = (count / total) * idf_val;
}
}
let magnitude: f32 = vec.iter().map(|v| v * v).sum::<f32>().sqrt();
if magnitude > 1e-9 {
for v in &mut vec {
*v /= magnitude;
}
}
vec
}
pub fn embed_tokens(&self, tokens: &[String]) -> Vec<f32> {
let mut vec = vec![0.0f32; self.dimension];
if self.vocab.is_empty() {
return vec;
}
let total = tokens.len() as f32;
if total == 0.0 {
return vec;
}
let mut tf_map: std::collections::HashMap<&str, f32> = std::collections::HashMap::new();
for tok in tokens {
*tf_map.entry(tok.as_str()).or_insert(0.0) += 1.0;
}
for (slot, (word, idf_val)) in vec.iter_mut().zip(self.vocab.iter().zip(self.idf.iter())) {
if let Some(&count) = tf_map.get(word.as_str()) {
*slot = (count / total) * idf_val;
}
}
let magnitude: f32 = vec.iter().map(|v| v * v).sum::<f32>().sqrt();
if magnitude > 1e-9 {
for v in &mut vec {
*v /= magnitude;
}
}
vec
}
fn from_persisted_state(state: TfIdfPersistedState) -> Option<Self> {
if state.schema_version != TFIDF_SCHEMA_VERSION {
tracing::warn!(
"Persisted TF-IDF schema version {} != current {}; discarding",
state.schema_version,
TFIDF_SCHEMA_VERSION
);
return None;
}
if state.dimension != crate::search::search::DEFAULT_EMBEDDING_DIMENSION {
tracing::warn!(
"Persisted TF-IDF dimension {} != expected {}; discarding",
state.dimension,
crate::search::search::DEFAULT_EMBEDDING_DIMENSION
);
return None;
}
if state.vocab.len() != state.idf.len() {
tracing::warn!(
"Persisted TF-IDF vocab/idf length mismatch ({} != {}); discarding",
state.vocab.len(),
state.idf.len()
);
return None;
}
Some(Self {
vocab: state.vocab,
idf: state.idf,
dimension: state.dimension,
pdg_nodes: state.pdg_nodes,
pdg_edges: state.pdg_edges,
pdg_fingerprint: state.pdg_fingerprint,
})
}
fn persisted_state(&self, pdg: &ProgramDependenceGraph) -> TfIdfPersistedState {
TfIdfPersistedState {
schema_version: TFIDF_SCHEMA_VERSION,
vocab: self.vocab.clone(),
idf: self.idf.clone(),
dimension: self.dimension,
pdg_nodes: pdg.node_count(),
pdg_edges: pdg.edge_count(),
pdg_fingerprint: pdg_search_fingerprint(pdg),
}
}
pub fn is_fresh(
&self,
pdg_node_count: usize,
pdg_edge_count: usize,
pdg_fingerprint: &str,
) -> bool {
self.pdg_nodes == pdg_node_count
&& self.pdg_edges == pdg_edge_count
&& !pdg_fingerprint.is_empty()
&& self.pdg_fingerprint == pdg_fingerprint
}
pub fn dimension(&self) -> usize {
self.dimension
}
fn storage_path(project_path: &Path) -> PathBuf {
project_path.join(".leindex").join("tfidf_embedder.bin")
}
pub fn load_from_storage(project_path: &Path) -> Result<Option<Self>> {
Self::load_from_artifact_path(&project_path.join(".leindex"))
}
pub(crate) fn load_from_artifact_path(storage_path: &Path) -> Result<Option<Self>> {
let path = storage_path.join("tfidf_embedder.bin");
if !path.exists() {
return Ok(None);
}
let bytes = std::fs::read(&path)
.with_context(|| format!("Failed to read persisted embedder: {}", path.display()))?;
let state: TfIdfPersistedState = bincode::deserialize(&bytes).with_context(|| {
format!(
"Failed to deserialize persisted embedder: {}",
path.display()
)
})?;
Ok(Self::from_persisted_state(state))
}
pub fn persist_to_storage(
&self,
project_path: &Path,
pdg: &ProgramDependenceGraph,
) -> Result<()> {
let path = Self::storage_path(project_path);
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent).with_context(|| {
format!("Failed to create embedder directory: {}", parent.display())
})?;
}
let payload = bincode::serialize(&self.persisted_state(pdg))
.context("Failed to serialize embedder")?;
std::fs::write(&path, payload)
.with_context(|| format!("Failed to persist embedder: {}", path.display()))
}
}
#[cfg(test)]
mod tfidf_persistence_tests {
use super::*;
fn valid_state() -> TfIdfPersistedState {
TfIdfPersistedState {
schema_version: TFIDF_SCHEMA_VERSION,
vocab: vec!["alpha".to_string()],
idf: vec![1.0],
dimension: crate::search::search::DEFAULT_EMBEDDING_DIMENSION,
pdg_nodes: 1,
pdg_edges: 0,
pdg_fingerprint: "fp".to_string(),
}
}
#[test]
fn rejects_unknown_schema_version() {
let mut state = valid_state();
state.schema_version = 0;
assert!(TfIdfEmbedder::from_persisted_state(state).is_none());
}
#[test]
fn rejects_dimension_mismatch() {
let mut state = valid_state();
state.dimension = 512;
assert!(TfIdfEmbedder::from_persisted_state(state).is_none());
}
#[test]
fn rejects_vocab_idf_length_mismatch() {
let mut state = valid_state();
state.vocab.push("beta".to_string());
assert!(TfIdfEmbedder::from_persisted_state(state).is_none());
}
#[test]
fn accepts_valid_state_and_preserves_fingerprint() {
let state = valid_state();
let loaded = TfIdfEmbedder::from_persisted_state(state).expect("valid state loads");
assert_eq!(
loaded.dimension,
crate::search::search::DEFAULT_EMBEDDING_DIMENSION
);
assert_eq!(loaded.pdg_fingerprint, "fp");
}
}