use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PackItem {
pub path: String,
pub lang: String,
pub name: String,
pub kind: String,
pub line_start: u32,
pub line_end: u32,
#[serde(skip_serializing_if = "Option::is_none")]
pub signature: Option<String>,
pub snippet_start: u32,
pub code: String,
pub reason: String,
pub score: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ContextPack {
pub task: String,
pub budget_tokens: u64,
pub used_tokens: u64,
pub truncated: bool,
pub items: Vec<PackItem>,
}
pub const CHARS_PER_TOKEN: u64 = 4;
pub fn est_tokens(chars: u64) -> u64 {
chars / CHARS_PER_TOKEN
}
pub fn tokenize(task: &str) -> Vec<String> {
let mut terms: Vec<String> = Vec::new();
let mut seen = std::collections::HashSet::new();
for raw in task.split(|c: char| !(c.is_alphanumeric() || c == '_')) {
if raw.is_empty() {
continue;
}
for tok in split_identifier(raw) {
if tok.len() < 2 || is_stopword(&tok) {
continue;
}
if seen.insert(tok.clone()) {
terms.push(tok);
}
}
}
terms
}
fn is_stopword(t: &str) -> bool {
matches!(
t,
"the"
| "a"
| "an"
| "of"
| "to"
| "in"
| "is"
| "for"
| "and"
| "or"
| "how"
| "where"
| "what"
| "does"
| "do"
| "with"
| "on"
| "by"
| "this"
| "that"
| "it"
| "be"
| "as"
| "at"
| "we"
| "i"
| "add"
| "fix"
| "use"
| "using"
| "make"
| "get"
| "set"
| "all"
| "when"
| "from"
| "into"
| "via"
| "can"
| "should"
| "code"
| "function"
| "method"
)
}
pub fn lexical_score(
name: &str,
kind: &str,
signature: Option<&str>,
container: Option<&str>,
path: &str,
terms: &[String],
) -> f32 {
if terms.is_empty() {
return 0.0;
}
let mut scratch = ScoreScratch::default();
let path_bonus = path_term_bonus(path, terms, &mut scratch);
lexical_score_with(
name,
kind,
signature,
container,
path_bonus,
terms,
&mut scratch,
)
}
#[derive(Default)]
pub struct ScoreScratch {
path_lower: String,
name_lower: String,
sig_lower: String,
name_tokens: Vec<String>,
cont_tokens: Vec<String>,
}
fn ascii_lower_into(s: &str, buf: &mut String) {
buf.clear();
buf.push_str(s);
buf.make_ascii_lowercase();
}
pub fn path_term_bonus(path: &str, terms: &[String], scratch: &mut ScoreScratch) -> f32 {
ascii_lower_into(path, &mut scratch.path_lower);
let mut bonus = 0.0f32;
for term in terms {
if scratch.path_lower.contains(term.as_str()) {
bonus += 2.0;
}
}
bonus
}
pub fn lexical_score_with(
name: &str,
kind: &str,
signature: Option<&str>,
container: Option<&str>,
path_bonus: f32,
terms: &[String],
scratch: &mut ScoreScratch,
) -> f32 {
if terms.is_empty() {
return 0.0;
}
ascii_lower_into(name, &mut scratch.name_lower);
let n_name = split_identifier_into(name, &mut scratch.name_tokens);
match signature {
Some(s) => ascii_lower_into(s, &mut scratch.sig_lower),
None => scratch.sig_lower.clear(),
}
let n_cont = match container {
Some(c) => split_identifier_into(c, &mut scratch.cont_tokens),
None => 0,
};
let name_lower = &scratch.name_lower;
let mut score = path_bonus;
for term in terms {
if name_lower == term {
score += 20.0;
} else if scratch.name_tokens[..n_name].iter().any(|t| t == term) {
score += 12.0;
} else if name_lower.contains(term.as_str()) {
score += 6.0;
}
if scratch.cont_tokens[..n_cont].iter().any(|t| t == term) {
score += 4.0;
}
if signature.is_some() && scratch.sig_lower.contains(term.as_str()) {
score += 3.0;
}
}
if score > 0.0 && is_priority_kind(kind) {
score += 2.0;
}
score
}
fn is_priority_kind(kind: &str) -> bool {
matches!(
kind,
"function"
| "method"
| "struct"
| "class"
| "trait"
| "interface"
| "enum"
| "type"
| "constructor"
| "module"
)
}
pub fn split_identifier(s: &str) -> Vec<String> {
let mut tokens = Vec::new();
let n = split_identifier_into(s, &mut tokens);
tokens.truncate(n);
tokens
}
fn split_identifier_into(s: &str, out: &mut Vec<String>) -> usize {
let mut n = 0usize;
let mut prev_lower = false;
let mut open = false;
for ch in s.chars() {
if ch == '_' || ch == '-' || ch == ' ' {
if open {
n += 1;
open = false;
}
prev_lower = false;
continue;
}
if ch.is_uppercase() && prev_lower && open {
n += 1;
open = false;
}
if !open {
if out.len() == n {
out.push(String::new());
} else {
out[n].clear();
}
open = true;
}
out[n].extend(ch.to_lowercase());
prev_lower = ch.is_lowercase() || ch.is_numeric();
}
if open {
n += 1;
}
n
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn tokenizes_and_drops_stopwords() {
let t = tokenize("How does the SegmentWriter flush to disk?");
assert!(t.contains(&"segment".to_string()));
assert!(t.contains(&"writer".to_string()));
assert!(t.contains(&"flush".to_string()));
assert!(t.contains(&"disk".to_string()));
assert!(!t.contains(&"the".to_string()));
assert!(!t.contains(&"how".to_string()));
}
#[test]
fn scores_name_matches_highest() {
let terms = tokenize("flush segment writer");
let exact = lexical_score("flush", "function", None, None, "src/a.rs", &terms);
let unrelated = lexical_score("zebra", "function", None, None, "src/a.rs", &terms);
assert!(exact > unrelated);
assert_eq!(unrelated, 0.0);
}
const IDENTS: &[&str] = &[
"",
"x",
"flush",
"loadConfig",
"load_config",
"kebab-case-name",
"with space",
"__leading",
"trailing__",
"a__b",
"HTTPServer",
"parseHTTP2Frame",
"v2Handler",
"snake_And_Camel",
"ÄÖÜ_grüß",
"İstanbul",
"ALLCAPS",
];
#[test]
fn split_into_matches_allocating_version_and_reuses_buffer() {
let mut buf: Vec<String> = Vec::new();
for s in IDENTS {
let want = split_identifier(s);
let n = split_identifier_into(s, &mut buf);
assert_eq!(n, want.len(), "token count for {s:?}");
assert_eq!(&buf[..n], &want[..], "tokens for {s:?}");
}
for s in IDENTS.iter().rev() {
let want = split_identifier(s);
let n = split_identifier_into(s, &mut buf);
assert_eq!(&buf[..n], &want[..], "tokens for {s:?} after reuse");
}
}
#[test]
fn scratch_reuse_does_not_change_scores() {
let terms = tokenize("write back page cache to disk");
type Case<'a> = (&'a str, &'a str, Option<&'a str>, Option<&'a str>, &'a str);
let cases: &[Case] = &[
("writeback", "function", None, None, "mm/page-writeback.c"),
(
"write_back_pages",
"function",
Some("int write_back_pages(struct page *p)"),
Some("PageCache"),
"fs/read_write.c",
),
("zebra", "struct", None, None, "drivers/zoo.c"),
(
"cache",
"method",
Some("void cache(void)"),
None,
"mm/cache.c",
),
(
"İstanbul",
"function",
None,
Some("ÄÖÜ_grüß"),
"i18n/ünicode.c",
),
("x", "field", Some("disk"), Some("page"), "a.c"),
("PageWriteback", "class", None, None, "include/linux/page.h"),
];
let mut shared = ScoreScratch::default();
for (name, kind, sig, cont, path) in cases {
let mut fresh = ScoreScratch::default();
let want = lexical_score_with(
name,
kind,
*sig,
*cont,
path_term_bonus(path, &terms, &mut fresh),
&terms,
&mut fresh,
);
let got = lexical_score_with(
name,
kind,
*sig,
*cont,
path_term_bonus(path, &terms, &mut shared),
&terms,
&mut shared,
);
assert_eq!(got, want, "score for {name:?} in {path:?}");
assert_eq!(
lexical_score(name, kind, *sig, *cont, path, &terms),
want,
"wrapper score for {name:?}"
);
}
}
#[test]
fn empty_and_none_signature_container_are_distinguished() {
let terms = tokenize("alpha beta");
let mut sh = ScoreScratch::default();
let _ = lexical_score_with(
"alpha_beta_gamma",
"function",
Some("fn alpha(beta: Beta) -> Gamma"),
Some("AlphaContainer"),
0.0,
&terms,
&mut sh,
);
let cases: &[(Option<&str>, Option<&str>)] = &[
(None, None),
(Some(""), None),
(None, Some("")),
(Some(""), Some("")),
(Some("alpha"), Some("beta")),
];
for (sig, cont) in cases {
let mut fresh = ScoreScratch::default();
let want = lexical_score_with("zzz", "other", *sig, *cont, 0.0, &terms, &mut fresh);
let got = lexical_score_with("zzz", "other", *sig, *cont, 0.0, &terms, &mut sh);
assert_eq!(got, want, "sig={sig:?} cont={cont:?}");
assert_eq!(
lexical_score("zzz", "other", *sig, *cont, "x.rs", &terms),
want,
"wrapper disagrees for sig={sig:?} cont={cont:?}"
);
}
}
#[test]
fn path_bonus_counts_each_matching_term_once() {
let terms = tokenize("page cache writeback");
let mut s = ScoreScratch::default();
assert_eq!(path_term_bonus("mm/nothing.c", &terms, &mut s), 0.0);
assert_eq!(path_term_bonus("mm/PAGE.c", &terms, &mut s), 2.0);
assert_eq!(path_term_bonus("mm/page-writeback.c", &terms, &mut s), 4.0);
assert!(lexical_score("zebra", "other", None, None, "mm/page.c", &terms) > 0.0);
}
}