scryer-engine 0.2.1

Tree-sitter and stack-graphs AST indexing engine for Scryer code intelligence
//! In-memory BM25F index with named, weighted fields.
//!
//! Each field's term frequency is length-normalised against that field's average length,
//! weighted, and summed before the standard BM25 saturation, then multiplied by IDF.
//! Ranking is deterministic: equal scores are ordered by document insertion order.

use std::collections::HashMap;

const K1: f32 = 1.2;
const B: f32 = 0.75;

#[derive(Debug, Clone, Copy)]
struct Posting {
    doc: u32,
    field: u8,
    tf: u16,
}

/// A ranked search hit. Resolve the payload with [`Bm25Index::doc`].
#[derive(Debug, Clone, PartialEq)]
pub struct Hit {
    pub doc_idx: usize,
    pub score: f32,
    /// Query terms that matched this document, in query order.
    pub matched_terms: Vec<String>,
}

/// Immutable BM25F index over documents of type `D`.
#[derive(Debug)]
pub struct Bm25Index<D> {
    field_weights: Vec<f32>,
    docs: Vec<D>,
    /// Flattened `[doc * num_fields + field]` token counts.
    field_lens: Vec<u32>,
    avg_lens: Vec<f32>,
    postings: HashMap<String, Vec<Posting>>,
}

/// Incremental builder for [`Bm25Index`].
#[derive(Debug)]
pub struct Bm25Builder<D> {
    index: Bm25Index<D>,
}

impl<D> Bm25Builder<D> {
    /// Creates a builder for the given per-field weights (field index = position).
    pub fn new(field_weights: &[f32]) -> Self {
        assert!(field_weights.len() <= u8::MAX as usize, "too many fields");
        Self {
            index: Bm25Index {
                field_weights: field_weights.to_vec(),
                docs: Vec::new(),
                field_lens: Vec::new(),
                avg_lens: vec![0.0; field_weights.len()],
                postings: HashMap::new(),
            },
        }
    }

    /// Adds a document. `fields[i]` holds the tokens of field `i`; missing trailing fields
    /// are treated as empty.
    pub fn add(&mut self, payload: D, fields: &[Vec<String>]) {
        let num_fields = self.index.field_weights.len();
        let doc = self.index.docs.len() as u32;
        self.index.docs.push(payload);
        for field in 0..num_fields {
            let tokens = fields.get(field).map(Vec::as_slice).unwrap_or(&[]);
            self.index.field_lens.push(tokens.len() as u32);
            let mut counts: Vec<(&str, u16)> = Vec::new();
            for token in tokens {
                match counts.iter_mut().find(|(t, _)| *t == token.as_str()) {
                    Some((_, c)) => *c = c.saturating_add(1),
                    None => counts.push((token.as_str(), 1)),
                }
            }
            for (term, tf) in counts {
                self.index
                    .postings
                    .entry(term.to_string())
                    .or_default()
                    .push(Posting {
                        doc,
                        field: field as u8,
                        tf,
                    });
            }
        }
    }

    pub fn build(mut self) -> Bm25Index<D> {
        let num_fields = self.index.field_weights.len();
        let num_docs = self.index.docs.len();
        if num_docs > 0 {
            for field in 0..num_fields {
                let total: u64 = (0..num_docs)
                    .map(|d| u64::from(self.index.field_lens[d * num_fields + field]))
                    .sum();
                self.index.avg_lens[field] = total as f32 / num_docs as f32;
            }
        }
        self.index
    }
}

impl<D> Bm25Index<D> {
    pub fn len(&self) -> usize {
        self.docs.len()
    }

    pub fn is_empty(&self) -> bool {
        self.docs.is_empty()
    }

    pub fn doc(&self, idx: usize) -> &D {
        &self.docs[idx]
    }

    pub fn docs(&self) -> &[D] {
        &self.docs
    }

    /// Returns every document with a positive score that passes `filter`, sorted by score
    /// descending and then by insertion order.
    pub fn search(&self, query: &[String], filter: impl Fn(&D) -> bool) -> Vec<Hit> {
        let num_fields = self.field_weights.len();
        let n = self.docs.len() as f32;
        let mut acc: HashMap<u32, (f32, Vec<usize>)> = HashMap::new();
        let mut allowed: HashMap<u32, bool> = HashMap::new();

        let mut seen_terms: Vec<&str> = Vec::new();
        for (term_idx, term) in query.iter().enumerate() {
            if seen_terms.contains(&term.as_str()) {
                continue;
            }
            seen_terms.push(term);
            let Some(postings) = self.postings.get(term) else {
                continue;
            };

            // Postings for one doc are contiguous (docs are added sequentially).
            let mut per_doc: Vec<(u32, f32)> = Vec::new();
            for p in postings {
                let len = self.field_lens[p.doc as usize * num_fields + p.field as usize] as f32;
                let avg = self.avg_lens[p.field as usize].max(f32::EPSILON);
                let norm = 1.0 - B + B * len / avg;
                let weighted = self.field_weights[p.field as usize] * f32::from(p.tf) / norm;
                match per_doc.last_mut() {
                    Some((d, tf)) if *d == p.doc => *tf += weighted,
                    _ => per_doc.push((p.doc, weighted)),
                }
            }

            let df = per_doc.len() as f32;
            let idf = (1.0 + (n - df + 0.5) / (df + 0.5)).ln();
            for (doc, tf) in per_doc {
                let ok = *allowed
                    .entry(doc)
                    .or_insert_with(|| filter(&self.docs[doc as usize]));
                if !ok {
                    continue;
                }
                let score = idf * tf * (K1 + 1.0) / (tf + K1);
                let entry = acc.entry(doc).or_insert((0.0, Vec::new()));
                entry.0 += score;
                entry.1.push(term_idx);
            }
        }

        let mut hits: Vec<Hit> = acc
            .into_iter()
            .filter(|(_, (score, _))| *score > 0.0)
            .map(|(doc, (score, terms))| Hit {
                doc_idx: doc as usize,
                score,
                matched_terms: terms.into_iter().map(|i| query[i].clone()).collect(),
            })
            .collect();
        hits.sort_by(|a, b| {
            b.score
                .total_cmp(&a.score)
                .then_with(|| a.doc_idx.cmp(&b.doc_idx))
        });
        hits
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    fn toks(s: &str) -> Vec<String> {
        s.split_whitespace().map(str::to_string).collect()
    }

    #[test]
    fn field_weights_change_ranking() {
        // doc 0 has "retry" in the low-weight field, doc 1 in the high-weight field.
        let mut b = Bm25Builder::new(&[3.0, 1.0]);
        b.add("body", &[toks("alpha"), toks("retry")]);
        b.add("name", &[toks("retry"), toks("alpha")]);
        let idx = b.build();
        let hits = idx.search(&toks("retry"), |_| true);
        assert_eq!(hits.len(), 2);
        assert_eq!(*idx.doc(hits[0].doc_idx), "name");
    }

    #[test]
    fn rare_term_outranks_common_term() {
        let mut b = Bm25Builder::new(&[1.0]);
        b.add(0, &[toks("common rare")]);
        b.add(1, &[toks("common other")]);
        b.add(2, &[toks("common")]);
        b.add(3, &[toks("common")]);
        let idx = b.build();
        let rare = idx.search(&toks("rare"), |_| true);
        let common = idx.search(&toks("common"), |_| true);
        assert!(rare[0].score > common[0].score);
        let both = idx.search(&toks("common rare"), |_| true);
        assert_eq!(both[0].doc_idx, 0);
    }

    #[test]
    fn equal_scores_are_ordered_by_insertion() {
        let mut b = Bm25Builder::new(&[1.0]);
        for i in 0..5 {
            b.add(i, &[toks("same tokens")]);
        }
        let idx = b.build();
        let hits = idx.search(&toks("same"), |_| true);
        let order: Vec<usize> = hits.iter().map(|h| h.doc_idx).collect();
        assert_eq!(order, vec![0, 1, 2, 3, 4]);
    }

    #[test]
    fn matched_terms_and_filter() {
        let mut b = Bm25Builder::new(&[1.0]);
        b.add(10, &[toks("retry backoff")]);
        b.add(20, &[toks("backoff")]);
        b.add(30, &[toks("parse config")]);
        let idx = b.build();
        let hits = idx.search(&toks("retry backoff missing"), |_| true);
        assert_eq!(hits.len(), 2);
        assert_eq!(hits[0].doc_idx, 0);
        assert_eq!(hits[0].matched_terms, toks("retry backoff"));
        assert_eq!(hits[1].matched_terms, toks("backoff"));

        let filtered = idx.search(&toks("backoff"), |d| *d != 10);
        assert_eq!(filtered.len(), 1);
        assert_eq!(*idx.doc(filtered[0].doc_idx), 20);
    }

    #[test]
    fn empty_index_and_query() {
        let idx: Bm25Index<u8> = Bm25Builder::new(&[1.0]).build();
        assert!(idx.search(&toks("x"), |_| true).is_empty());
        let mut b = Bm25Builder::new(&[1.0]);
        b.add(1u8, &[toks("x")]);
        assert!(b.build().search(&[], |_| true).is_empty());
    }
}