1use serde::{Deserialize, Serialize};
8
9pub const BM25_K1: f64 = 1.5;
11pub const BM25_B: f64 = 0.75;
13pub const WEIGHT_NAME: u32 = 3;
15pub const WEIGHT_PATH: u32 = 2;
17pub const WEIGHT_DOC: u32 = 2;
19pub const WEIGHT_BODY: u32 = 1;
21pub const MIN_SUBTOKEN_LEN: usize = 2;
23
24pub fn subtokens(text: &str) -> Vec<String> {
29 let mut out = Vec::new();
30 for_each_lex_subtoken(text, |start, end| {
31 if end.saturating_sub(start) < MIN_SUBTOKEN_LEN {
32 return;
33 }
34 let tok: String = text.as_bytes()[start..end]
35 .iter()
36 .map(|b| lex_lower_byte(*b) as char)
37 .collect();
38 out.push(tok);
39 });
40 out
41}
42
43fn lex_lower_byte(c: u8) -> u8 {
45 if c.is_ascii_uppercase() {
46 c - b'A' + b'a'
47 } else {
48 c
49 }
50}
51
52fn lex_upper_opens_token(text: &[u8], k: usize, prev_upper: bool) -> bool {
54 let next = if k + 1 < text.len() { text[k + 1] } else { 0 };
55 !prev_upper || next.is_ascii_lowercase()
56}
57
58fn for_each_lex_subtoken(text: &str, mut emit: impl FnMut(usize, usize)) {
60 let bytes = text.as_bytes();
61 let mut tok_start: Option<usize> = None;
62 let mut prev_upper = false;
63 for k in 0..bytes.len() {
64 let c = bytes[k];
65 let upper = c.is_ascii_uppercase();
66 let lower = c.is_ascii_lowercase();
67 let digit = c.is_ascii_digit();
68 if !upper && !lower && !digit {
69 if let Some(s) = tok_start.take() {
70 emit(s, k);
71 }
72 prev_upper = false;
73 continue;
74 }
75 if upper {
76 if let Some(s) = tok_start {
77 if lex_upper_opens_token(bytes, k, prev_upper) {
78 emit(s, k);
79 tok_start = Some(k);
80 }
81 }
82 }
83 if tok_start.is_none() {
84 tok_start = Some(k);
85 }
86 prev_upper = upper;
87 }
88 if let Some(s) = tok_start {
89 emit(s, bytes.len());
90 }
91}
92
93#[derive(Debug, Clone)]
95pub struct LexField {
97 pub text: String,
98 pub weight: u32,
99}
100
101#[derive(Debug, Clone)]
103pub struct LexDoc {
105 pub id: String,
106 pub fields: Vec<LexField>,
107}
108
109impl LexDoc {
111 pub fn from_parts(
113 id: impl Into<String>,
114 name: &str,
115 path: &str,
116 doc: &str,
117 body: &str,
118 ) -> Self {
119 LexDoc {
120 id: id.into(),
121 fields: vec![
122 LexField {
123 text: name.to_string(),
124 weight: WEIGHT_NAME,
125 },
126 LexField {
127 text: path.to_string(),
128 weight: WEIGHT_PATH,
129 },
130 LexField {
131 text: doc.to_string(),
132 weight: WEIGHT_DOC,
133 },
134 LexField {
135 text: body.to_string(),
136 weight: WEIGHT_BODY,
137 },
138 ],
139 }
140 }
141}
142
143pub fn bm25_scores(query: &str, docs: &[LexDoc]) -> Vec<f64> {
146 let q_toks = subtokens(query);
147 let n = docs.len();
148 if q_toks.is_empty() || n == 0 {
149 return vec![0.0; n];
150 }
151
152 let mut unique: Vec<String> = Vec::new();
154 let mut q_index: Vec<usize> = Vec::with_capacity(q_toks.len());
155 for t in &q_toks {
156 if let Some(i) = unique.iter().position(|u| u == t) {
157 q_index.push(i);
158 } else {
159 q_index.push(unique.len());
160 unique.push(t.clone());
161 }
162 }
163 let u_count = unique.len();
164
165 let mut dl = vec![0u32; n];
166 let mut tf = vec![0u32; n * u_count];
167 for (i, doc) in docs.iter().enumerate() {
168 for field in &doc.fields {
169 if field.weight == 0 {
170 continue;
171 }
172 for tok in subtokens(&field.text) {
173 dl[i] = dl[i].saturating_add(field.weight);
174 if let Some(u) = unique.iter().position(|t| t == &tok) {
175 tf[i * u_count + u] = tf[i * u_count + u].saturating_add(field.weight);
176 }
177 }
178 }
179 }
180
181 let avgdl = if n == 0 {
182 1.0
183 } else {
184 dl.iter().map(|d| *d as f64).sum::<f64>() / n as f64
185 };
186 let avgdl = if avgdl > 0.0 { avgdl } else { 1.0 };
187
188 let mut df = vec![0u32; u_count];
189 for i in 0..n {
190 for u in 0..u_count {
191 if tf[i * u_count + u] > 0 {
192 df[u] += 1;
193 }
194 }
195 }
196
197 let mut scores = vec![0.0; n];
198 for i in 0..n {
199 let mut sc = 0.0;
200 for &u in &q_index {
202 let term_tf = tf[i * u_count + u] as f64;
203 if term_tf == 0.0 {
204 continue;
205 }
206 let n_df = df[u] as f64;
207 let idf = ((n as f64 - n_df + 0.5) / (n_df + 0.5) + 1.0).ln();
208 let denom = term_tf + BM25_K1 * (1.0 - BM25_B + BM25_B * (dl[i] as f64) / avgdl);
209 sc += idf * (term_tf * (BM25_K1 + 1.0)) / denom;
210 }
211 scores[i] = sc;
212 }
213 scores
214}
215
216#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
220pub struct Bm25CorpusStats {
222 pub n: u32,
223 pub avgdl: f64,
224 pub df: std::collections::BTreeMap<String, u32>,
226 pub dl: std::collections::BTreeMap<String, u32>,
228}
229
230impl Bm25CorpusStats {
232 pub fn from_docs(docs: &[LexDoc]) -> Self {
236 let mut df: std::collections::BTreeMap<String, u32> = std::collections::BTreeMap::new();
237 let mut dl: std::collections::BTreeMap<String, u32> = std::collections::BTreeMap::new();
238 for doc in docs {
239 let mut seen: std::collections::BTreeSet<String> = std::collections::BTreeSet::new();
240 let mut len = 0u32;
241 for field in &doc.fields {
242 if field.weight == 0 {
243 continue;
244 }
245 for tok in subtokens(&field.text) {
246 len = len.saturating_add(field.weight);
247 seen.insert(tok);
248 }
249 }
250 dl.insert(doc.id.clone(), len);
251 for tok in seen {
252 *df.entry(tok).or_insert(0) += 1;
253 }
254 }
255 let n = docs.len() as u32;
256 let avgdl = if n == 0 {
257 1.0
258 } else {
259 let sum: f64 = dl.values().map(|d| *d as f64).sum();
260 let a = sum / n as f64;
261 if a > 0.0 {
262 a
263 } else {
264 1.0
265 }
266 };
267 Bm25CorpusStats { n, avgdl, df, dl }
268 }
269}
270
271pub fn bm25_scores_with_stats(query: &str, docs: &[LexDoc], stats: &Bm25CorpusStats) -> Vec<f64> {
276 let q_toks = subtokens(query);
277 let n_docs = docs.len();
278 if q_toks.is_empty() || n_docs == 0 || stats.n == 0 {
279 return vec![0.0; n_docs];
280 }
281 let mut unique: Vec<String> = Vec::new();
282 let mut q_index: Vec<usize> = Vec::with_capacity(q_toks.len());
283 for t in &q_toks {
284 if let Some(i) = unique.iter().position(|u| u == t) {
285 q_index.push(i);
286 } else {
287 q_index.push(unique.len());
288 unique.push(t.clone());
289 }
290 }
291 let u_count = unique.len();
292 let mut dl = vec![0u32; n_docs];
293 let mut tf = vec![0u32; n_docs * u_count];
294 for (i, doc) in docs.iter().enumerate() {
295 dl[i] = stats.dl.get(&doc.id).copied().unwrap_or(0);
296 for field in &doc.fields {
297 if field.weight == 0 {
298 continue;
299 }
300 for tok in subtokens(&field.text) {
301 if !stats.dl.contains_key(&doc.id) {
302 dl[i] = dl[i].saturating_add(field.weight);
303 }
304 if let Some(u) = unique.iter().position(|t| t == &tok) {
305 tf[i * u_count + u] = tf[i * u_count + u].saturating_add(field.weight);
306 }
307 }
308 }
309 }
310 let avgdl = if stats.avgdl > 0.0 { stats.avgdl } else { 1.0 };
311 let n = stats.n as f64;
312 let mut scores = vec![0.0; n_docs];
313 for i in 0..n_docs {
314 let mut sc = 0.0;
315 let doc_len = if dl[i] == 0 { 1.0 } else { dl[i] as f64 };
316 for &u in &q_index {
317 let term_tf = tf[i * u_count + u] as f64;
318 if term_tf == 0.0 {
319 continue;
320 }
321 let n_df = stats.df.get(&unique[u]).copied().unwrap_or(0) as f64;
322 let idf = ((n - n_df + 0.5) / (n_df + 0.5) + 1.0).ln();
323 let denom = term_tf + BM25_K1 * (1.0 - BM25_B + BM25_B * doc_len / avgdl);
324 sc += idf * (term_tf * (BM25_K1 + 1.0)) / denom;
325 }
326 scores[i] = sc;
327 }
328 scores
329}
330
331pub fn bm25_rank(query: &str, docs: &[LexDoc]) -> Vec<(String, f64)> {
334 let scores = bm25_scores(query, docs);
335 let mut pairs: Vec<(String, f64)> = docs
336 .iter()
337 .zip(scores)
338 .map(|(d, s)| (d.id.clone(), s))
339 .collect();
340 pairs.sort_by(|a, b| {
341 b.1.partial_cmp(&a.1)
342 .unwrap_or(std::cmp::Ordering::Equal)
343 .then_with(|| a.0.cmp(&b.0))
344 });
345 pairs
346}
347
348#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
351#[serde(rename_all = "snake_case")]
352pub enum QueryShape {
354 Identifier,
355 SymbolLike,
356 PathLike,
357 StackTrace,
358 ErrorMessage,
359 Conceptual,
360 Architecture,
361 Flow,
362 State,
363 Impact,
364}
365
366impl QueryShape {
368 pub fn as_str(self) -> &'static str {
370 match self {
371 QueryShape::Identifier => "identifier",
372 QueryShape::SymbolLike => "symbol_like",
373 QueryShape::PathLike => "path_like",
374 QueryShape::StackTrace => "stack_trace",
375 QueryShape::ErrorMessage => "error_message",
376 QueryShape::Conceptual => "conceptual",
377 QueryShape::Architecture => "architecture",
378 QueryShape::Flow => "flow",
379 QueryShape::State => "state",
380 QueryShape::Impact => "impact",
381 }
382 }
383}
384
385#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
387pub struct QueryLocus {
389 pub path: String,
390 pub line: u32,
391 pub symbol: Option<String>,
392}
393
394#[derive(Debug, Clone, PartialEq, Eq)]
396pub struct RetrievalPlan {
398 pub shape: QueryShape,
399 pub prefer_exact_name: bool,
400 pub prefer_bm25: bool,
401 pub prefer_atlas: bool,
402 pub prefer_locus: bool,
403 pub loci: Vec<QueryLocus>,
404}
405
406pub fn classify_query(query: &str) -> QueryShape {
409 route_query(query).shape
410}
411
412pub fn route_query(query: &str) -> RetrievalPlan {
415 let trimmed = query.trim();
416 let loci = extract_loci(trimmed);
417 let specific_frames = count_specific_frames(trimmed);
418 let generic_frames = loci.len();
419 if specific_frames > 0 || generic_frames >= 2 {
420 return plan(QueryShape::StackTrace, false, false, false, true, loci);
421 }
422 if looks_error_message(trimmed) {
423 return plan(
424 QueryShape::ErrorMessage,
425 false,
426 generic_frames > 0,
427 false,
428 generic_frames > 0,
429 loci,
430 );
431 }
432 if looks_path_query(trimmed) {
433 return plan(QueryShape::PathLike, true, false, false, generic_frames > 0, loci);
434 }
435 if looks_symbol_like(trimmed) {
436 return plan(QueryShape::SymbolLike, true, false, false, false, loci);
437 }
438 if looks_identifier(trimmed) {
439 return plan(QueryShape::Identifier, true, false, false, false, loci);
440 }
441 let lower = trimmed.to_ascii_lowercase();
442 if has_any(&lower, &["architecture", "system atlas", "component", "trust boundary", "deployment"])
443 {
444 return plan(QueryShape::Architecture, false, true, true, false, loci);
445 }
446 if has_any(&lower, &["what happens", "call flow", "causal flow"])
447 || (lower.contains("flow") && (lower.contains('?') || lower.contains("when")))
448 {
449 return plan(QueryShape::Flow, false, true, false, false, loci);
450 }
451 if has_any(
452 &lower,
453 &[
454 "state authority",
455 "source of truth",
456 "who writes",
457 "invalidat",
458 "cache owner",
459 ],
460 ) {
461 return plan(QueryShape::State, false, true, false, false, loci);
462 }
463 if has_any(
464 &lower,
465 &[
466 "blast radius",
467 "who calls",
468 "what breaks",
469 "impact of",
470 "if i change",
471 ],
472 ) {
473 return plan(QueryShape::Impact, false, true, false, false, loci);
474 }
475 plan(QueryShape::Conceptual, false, true, false, false, loci)
476}
477
478fn plan(
480 shape: QueryShape,
481 prefer_exact_name: bool,
482 prefer_bm25: bool,
483 prefer_atlas: bool,
484 prefer_locus: bool,
485 loci: Vec<QueryLocus>,
486) -> RetrievalPlan {
487 RetrievalPlan {
488 shape,
489 prefer_exact_name,
490 prefer_bm25,
491 prefer_atlas,
492 prefer_locus,
493 loci,
494 }
495}
496
497fn has_any(hay: &str, needles: &[&str]) -> bool {
499 needles.iter().any(|n| hay.contains(n))
500}
501
502fn looks_identifier(q: &str) -> bool {
504 let q = q.trim();
505 if q.is_empty() || q.contains(char::is_whitespace) || q.contains('/') || q.contains('\\') {
506 return false;
507 }
508 if q.contains("::") || q.contains('.') {
509 return false;
510 }
511 q.chars()
512 .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-')
513 && q.chars().any(|c| c.is_ascii_alphabetic())
514}
515
516fn looks_symbol_like(q: &str) -> bool {
518 let q = q.trim();
519 if q.is_empty() || q.contains(char::is_whitespace) || q.contains('/') || q.contains('\\') {
520 return false;
521 }
522 (q.contains("::") || q.contains('.'))
523 && q.chars()
524 .all(|c| c.is_ascii_alphanumeric() || matches!(c, '_' | ':' | '.' | '#'))
525}
526
527fn looks_path_query(q: &str) -> bool {
529 let q = q.trim();
530 if q.contains(char::is_whitespace) {
531 return false;
532 }
533 let slash = q.contains('/') || q.contains('\\');
534 let ext = q.rsplit_once('.').map(|(_, e)| {
535 !e.is_empty()
536 && e.len() <= 5
537 && e.chars().all(|c| c.is_ascii_alphanumeric())
538 });
539 slash && ext.unwrap_or(false) && !q.contains("://")
540}
541
542fn looks_error_message(q: &str) -> bool {
544 let lower = q.to_ascii_lowercase();
545 has_any(
546 &lower,
547 &[
548 "error:",
549 "exception",
550 "panic!",
551 "fatal:",
552 "failed to",
553 "undefined is not",
554 "cannot find",
555 "typeerror",
556 "nullpointer",
557 ],
558 )
559}
560
561fn count_specific_frames(q: &str) -> usize {
563 let mut n = 0;
564 for line in q.lines() {
565 let t = line.trim();
566 let lower = t.to_ascii_lowercase();
567 if lower.contains("file \"") && lower.contains(", line ") {
568 n += 1;
569 continue;
570 }
571 if t.starts_with("at ") && (t.contains('(') || t.contains(".java:")) {
572 n += 1;
573 continue;
574 }
575 if t.contains(" --> ") && t.contains(".rs:") {
576 n += 1;
577 continue;
578 }
579 if t.starts_with('#') && t.contains("0x") {
580 n += 1;
581 }
582 }
583 n
584}
585
586pub fn extract_loci(query: &str) -> Vec<QueryLocus> {
589 let mut out = Vec::new();
590 for line in query.lines() {
591 if let Some(loc) = locus_from_line(line) {
592 out.push(loc);
593 }
594 }
595 out
596}
597
598pub fn path_matches_locus(indexed: &str, locus: &str) -> bool {
602 fn norm(p: &str) -> String {
603 p.trim().trim_start_matches("./").replace('\\', "/")
604 }
605 let a = norm(indexed);
606 let b = norm(locus);
607 if a.is_empty() || b.is_empty() {
608 return false;
609 }
610 a == b || a.ends_with(&format!("/{b}")) || b.ends_with(&format!("/{a}"))
611}
612
613fn locus_from_line(line: &str) -> Option<QueryLocus> {
615 let t = line.trim();
616 if let Some(rest) = t.find("File \"").or_else(|| t.find("file \"")) {
618 let after = &t[rest + 6..];
619 if let Some(end) = after.find('"') {
620 let path = after[..end].to_string();
621 let tail = &after[end + 1..];
622 let line_no = parse_after(tail, "line ")?;
623 let symbol = tail
624 .rsplit("in ")
625 .next()
626 .map(|s| s.trim().trim_end_matches(',').to_string())
627 .filter(|s| {
628 !s.is_empty() && s.chars().all(|c| c.is_ascii_alphanumeric() || c == '_')
629 });
630 return Some(QueryLocus {
631 path,
632 line: line_no,
633 symbol,
634 });
635 }
636 }
637 if let Some(idx) = t.find("--> ") {
639 return path_line_from(&t[idx + 4..]);
640 }
641 path_line_from(t)
642}
643
644fn path_line_from(t: &str) -> Option<QueryLocus> {
646 let token = t
647 .split_whitespace()
648 .rev()
649 .find(|s| s.contains('.') && s.contains(':'))?;
650 let token = token.trim_end_matches([')', ',', ']']);
651 if token.contains("://") {
652 return None;
653 }
654 let mut parts: Vec<&str> = token.rsplitn(3, ':').collect();
655 parts.reverse();
656 let (path, line_s) = match parts.as_slice() {
657 [path, line] if line.chars().all(|c| c.is_ascii_digit()) => (*path, *line),
658 [path, line, col]
659 if line.chars().all(|c| c.is_ascii_digit()) && col.chars().all(|c| c.is_ascii_digit()) =>
660 {
661 (*path, *line)
662 }
663 _ => return None,
664 };
665 if !looks_source_path(path) {
666 return None;
667 }
668 let line: u32 = line_s.parse().ok()?;
669 Some(QueryLocus {
670 path: path.to_string(),
671 line,
672 symbol: None,
673 })
674}
675
676fn looks_source_path(path: &str) -> bool {
678 let ext = path.rsplit_once('.').map(|(_, e)| e).unwrap_or("");
679 !ext.is_empty()
680 && ext.len() <= 5
681 && ext.chars().all(|c| c.is_ascii_alphanumeric())
682 && !path.contains("://")
683}
684
685fn parse_after(hay: &str, key: &str) -> Option<u32> {
687 let i = hay.find(key)?;
688 let rest = hay[i + key.len()..].trim_start();
689 let num: String = rest.chars().take_while(|c| c.is_ascii_digit()).collect();
690 num.parse().ok()
691}
692
693#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, Default)]
696#[serde(rename_all = "kebab-case")]
697pub enum RankingArm {
699 #[default]
701 ProductionBlended,
702 LexicalThenGraph,
703 QueryRouted,
704 NoGraph,
705 NoLexical,
706 NoSemantic,
707 NoCochange,
708 NoTaskPpr,
709 NoGlobalPpr,
710}
711
712impl RankingArm {
714 pub fn as_str(self) -> &'static str {
716 match self {
717 RankingArm::ProductionBlended => "production-blended",
718 RankingArm::LexicalThenGraph => "lexical-then-graph",
719 RankingArm::QueryRouted => "query-routed",
720 RankingArm::NoGraph => "no-graph",
721 RankingArm::NoLexical => "no-lexical",
722 RankingArm::NoSemantic => "no-semantic",
723 RankingArm::NoCochange => "no-cochange",
724 RankingArm::NoTaskPpr => "no-task-ppr",
725 RankingArm::NoGlobalPpr => "no-global-ppr",
726 }
727 }
728
729 pub fn all() -> &'static [RankingArm] {
731 &[
732 RankingArm::ProductionBlended,
733 RankingArm::LexicalThenGraph,
734 RankingArm::QueryRouted,
735 RankingArm::NoGraph,
736 RankingArm::NoLexical,
737 RankingArm::NoSemantic,
738 RankingArm::NoCochange,
739 RankingArm::NoTaskPpr,
740 RankingArm::NoGlobalPpr,
741 ]
742 }
743
744 pub fn parse(s: &str) -> Option<Self> {
746 Self::all().iter().copied().find(|a| a.as_str() == s)
747 }
748}
749
750#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
752pub struct RelevanceHit {
754 pub id: String,
755 pub bm25: f64,
756 pub exact_anchor: bool,
757 pub shape: QueryShape,
758}
759
760pub const QUERY_MENTION_MAX_RAW: usize = 16;
762
763#[derive(Debug, Clone, PartialEq, Eq)]
765pub struct QueryMention {
767 pub segments: Vec<String>,
768 pub is_path: bool,
769 pub backticked: bool,
770}
771
772fn is_ident_char(c: char) -> bool {
774 c.is_ascii_alphanumeric() || c == '_'
775}
776
777fn is_mention_token_char(c: char) -> bool {
779 is_ident_char(c) || c == '.' || c == '/' || c == '-'
780}
781
782pub fn extract_query_mentions(task: &str) -> Vec<QueryMention> {
786 let bytes: Vec<char> = task.chars().collect();
787 let mut raw = Vec::new();
788 let mut i = 0usize;
789 while i < bytes.len() && raw.len() < QUERY_MENTION_MAX_RAW {
790 if !is_mention_token_char(bytes[i]) {
791 i += 1;
792 continue;
793 }
794 let start = i;
795 while i < bytes.len() && is_mention_token_char(bytes[i]) {
796 i += 1;
797 }
798 let mut tok: String = bytes[start..i].iter().collect();
799 let backticked = start > 0
800 && bytes[start - 1] == '`'
801 && i < bytes.len()
802 && bytes[i] == '`';
803 while tok.ends_with('.') || tok.ends_with('/') || tok.ends_with('-') {
804 tok.pop();
805 }
806 while tok.starts_with('.') || tok.starts_with('/') || tok.starts_with('-') {
807 tok.remove(0);
808 }
809 if tok.len() < 3 || tok.len() > 200 {
810 continue;
811 }
812 let has_slash = tok.contains('/');
813 let has_dot = tok.contains('.');
814 if !has_slash && !has_dot && !backticked {
815 continue;
816 }
817 let joiner = if has_slash { '/' } else { '.' };
818 let mut segments: Vec<String> = Vec::new();
819 let mut malformed = false;
820 let mut max_seg = 0usize;
821 for seg in tok.split(joiner) {
822 if seg.is_empty() {
823 if has_slash {
824 continue;
825 }
826 malformed = true;
827 break;
828 }
829 max_seg = max_seg.max(seg.len());
830 segments.push(seg.to_string());
831 }
832 if malformed || segments.is_empty() {
833 continue;
834 }
835 if !has_slash && !backticked {
836 if segments.len() < 2 || max_seg < 3 {
837 continue;
838 }
839 if segments
840 .iter()
841 .all(|s| s.chars().all(|c| c.is_ascii_digit()))
842 {
843 continue;
844 }
845 }
846 if has_slash && segments.len() > 3 {
847 segments = segments.split_off(segments.len() - 3);
848 }
849 raw.push(QueryMention {
850 segments,
851 is_path: has_slash,
852 backticked,
853 });
854 }
855 raw
856}
857
858fn path_basename(path: &str) -> &str {
860 path.rsplit(['/', '\\']).next().unwrap_or(path)
861}
862
863fn strip_ext(name: &str) -> &str {
865 match name.rfind('.') {
866 Some(i) if i > 0 => &name[..i],
867 _ => name,
868 }
869}
870
871fn path_suffix_matches(path: &str, segments: &[String]) -> bool {
875 if segments.is_empty() {
876 return false;
877 }
878 let mut comps: Vec<&str> = path
879 .split(['/', '\\'])
880 .filter(|s| !s.is_empty())
881 .collect();
882 if comps.is_empty() {
883 return false;
884 }
885 let last = *comps.last().unwrap();
886 let want_last = segments.last().unwrap().as_str();
887 if last != want_last && strip_ext(last) != want_last {
888 return false;
889 }
890 comps.pop();
891 let earlier = &segments[..segments.len() - 1];
892 if earlier.len() > comps.len() {
893 return false;
894 }
895 let tail = &comps[comps.len() - earlier.len()..];
896 tail.iter().zip(earlier.iter()).all(|(c, s)| *c == s.as_str())
897}
898
899pub fn mention_matches_doc(name: &str, path: &str, mention: &QueryMention) -> bool {
902 let joined = mention.segments.join(if mention.is_path { "/" } else { "." });
903 if !name.is_empty() {
904 if name == joined || name.eq_ignore_ascii_case(&joined) {
905 return true;
906 }
907 if let Some(last) = mention.segments.last() {
908 if name == last.as_str()
909 || name.ends_with(&format!(".{last}"))
910 || name.ends_with(&format!("::{last}"))
911 {
912 return true;
913 }
914 }
915 }
916 if path.is_empty() {
917 return false;
918 }
919 if path_suffix_matches(path, &mention.segments) {
920 return true;
921 }
922 for suffix_len in 1..mention.segments.len() {
923 let suffix = &mention.segments[mention.segments.len() - suffix_len..];
924 if path_suffix_matches(path, suffix) {
925 return true;
926 }
927 }
928 if let Some(last) = mention.segments.last() {
929 let base = path_basename(path);
930 if base == last.as_str() || strip_ext(base) == last.as_str() {
931 return true;
932 }
933 }
934 false
935}
936
937pub fn is_exact_anchor(query: &str, name: &str) -> bool {
941 let q = query.trim();
942 if q.is_empty() || name.is_empty() {
943 return false;
944 }
945 if q.eq_ignore_ascii_case(name) {
946 return true;
947 }
948 let q_last = q.rsplit(['/', '\\', '.', ':']).next().unwrap_or(q);
949 if q_last.eq_ignore_ascii_case(name) {
950 return true;
951 }
952 let q_toks = subtokens(q);
953 let n_toks = subtokens(name);
954 !q_toks.is_empty() && q_toks == n_toks
955}
956
957pub fn relevance_hits(query: &str, docs: &[LexDoc]) -> Vec<RelevanceHit> {
962 relevance_hits_from_scores(query, docs, bm25_scores(query, docs))
963}
964
965pub fn relevance_hits_with_stats(
969 query: &str,
970 docs: &[LexDoc],
971 stats: &Bm25CorpusStats,
972) -> Vec<RelevanceHit> {
973 relevance_hits_from_scores(query, docs, bm25_scores_with_stats(query, docs, stats))
974}
975
976fn relevance_hits_from_scores(
978 query: &str,
979 docs: &[LexDoc],
980 scores: Vec<f64>,
981) -> Vec<RelevanceHit> {
982 let shape = classify_query(query);
983 let mentions = extract_query_mentions(query);
984 let mut hits: Vec<RelevanceHit> = docs
985 .iter()
986 .zip(scores)
987 .map(|(d, bm25)| {
988 let name = anchor_name(d);
989 let path = path_field(d);
990 let named = is_exact_anchor(query, &name)
991 || mentions
992 .iter()
993 .any(|m| mention_matches_doc(&name, &path, m));
994 RelevanceHit {
995 id: d.id.clone(),
996 bm25,
997 exact_anchor: named,
998 shape,
999 }
1000 })
1001 .collect();
1002 hits.sort_by(|a, b| {
1003 b.exact_anchor
1004 .cmp(&a.exact_anchor)
1005 .then_with(|| {
1006 b.bm25
1007 .partial_cmp(&a.bm25)
1008 .unwrap_or(std::cmp::Ordering::Equal)
1009 })
1010 .then_with(|| a.id.cmp(&b.id))
1011 });
1012 hits
1013}
1014
1015fn anchor_name(doc: &LexDoc) -> String {
1017 doc.fields
1018 .iter()
1019 .find(|f| f.weight == WEIGHT_NAME)
1020 .map(|f| f.text.clone())
1021 .unwrap_or_default()
1022}
1023
1024fn path_field(doc: &LexDoc) -> String {
1026 doc.fields
1027 .iter()
1028 .find(|f| f.weight == WEIGHT_PATH)
1029 .map(|f| f.text.clone())
1030 .unwrap_or_default()
1031}
1032
1033pub fn ranking_arm_ids() -> Vec<&'static str> {
1037 RankingArm::all().iter().map(|a| a.as_str()).collect()
1038}
1039
1040#[cfg(test)]
1041mod tests {
1042 use super::*;
1043
1044 #[test]
1045 fn acronym_and_camel_splits_match_ripwire_rule() {
1047 assert_eq!(subtokens("MCP"), vec!["mcp"]);
1048 assert_eq!(subtokens("HTTPServer"), vec!["http", "server"]);
1049 assert_eq!(
1050 subtokens("updateCollisionPositionVelocity"),
1051 vec!["update", "collision", "position", "velocity"]
1052 );
1053 assert_eq!(subtokens("_max_speed"), vec!["max", "speed"]);
1054 assert_eq!(subtokens("IOError"), vec!["io", "error"]);
1055 assert_eq!(subtokens("XMLHttpRequest"), vec!["xml", "http", "request"]);
1056 assert_eq!(subtokens("a"), Vec::<String>::new()); assert_eq!(subtokens("handleList"), vec!["handle", "list"]);
1058 }
1059
1060 #[test]
1061 fn bm25_is_deterministic_and_name_weighted() {
1063 let docs = vec![
1064 LexDoc::from_parts("b", "unrelated", "src/other.py", "", "noise"),
1065 LexDoc::from_parts(
1066 "a",
1067 "handleList",
1068 "src/server.ts",
1069 "Fetch the user rows",
1070 "db.users.findMany",
1071 ),
1072 ];
1073 let s1 = bm25_scores("handleList", &docs);
1074 let s2 = bm25_scores("handleList", &docs);
1075 assert_eq!(s1, s2);
1076 assert!(s1[1] > s1[0], "name match must outrank unrelated: {s1:?}");
1077 let ranked = bm25_rank("handleList", &docs);
1078 assert_eq!(ranked[0].0, "a");
1079 assert_eq!(bm25_scores("", &docs), vec![0.0, 0.0]);
1081 }
1082
1083 #[test]
1084 fn query_shape_is_conservative() {
1086 assert_eq!(classify_query("handleList"), QueryShape::Identifier);
1087 assert_eq!(classify_query("Foo::bar"), QueryShape::SymbolLike);
1088 assert_eq!(classify_query("src/server.ts"), QueryShape::PathLike);
1089 assert_eq!(
1090 classify_query("https://example.com:8080/docs"),
1091 QueryShape::Conceptual
1092 );
1093 assert_eq!(
1094 classify_query("see Type.py:12 in the docs"),
1095 QueryShape::Conceptual
1096 ); let py = r#"Traceback (most recent call last):
1098 File "src/server.ts", line 12, in handleList
1099 db.users.findMany()
1100"#;
1101 assert_eq!(classify_query(py), QueryShape::StackTrace);
1102 let plan = route_query(py);
1103 assert!(plan.prefer_locus);
1104 assert!(!plan.loci.is_empty());
1105 assert_eq!(
1106 classify_query("how does authentication work?"),
1107 QueryShape::Conceptual
1108 );
1109 assert_eq!(
1110 classify_query("what is the system architecture of billing?"),
1111 QueryShape::Architecture
1112 );
1113 assert_eq!(
1114 classify_query("what happens when login fails?"),
1115 QueryShape::Flow
1116 );
1117 assert_eq!(
1118 classify_query("who writes SessionStore?"),
1119 QueryShape::State
1120 );
1121 assert_eq!(
1122 classify_query("blast radius if I change parseToken"),
1123 QueryShape::Impact
1124 );
1125 assert_eq!(
1126 classify_query("TypeError: cannot read property of undefined"),
1127 QueryShape::ErrorMessage
1128 );
1129 assert_eq!(RankingArm::default(), RankingArm::ProductionBlended);
1130 assert_eq!(ranking_arm_ids().len(), RankingArm::all().len());
1131 }
1132
1133 #[test]
1134 fn exact_anchor_is_not_mixed_into_bm25() {
1136 assert!(is_exact_anchor("handleList", "handleList"));
1137 assert!(is_exact_anchor("HTTPServer", "http_server"));
1138 assert!(!is_exact_anchor("how does login work", "handleList"));
1139 let docs = vec![
1140 LexDoc::from_parts("noise", "loginHelper", "src/a.ts", "handles login flow", ""),
1141 LexDoc::from_parts("hit", "handleList", "src/b.ts", "", ""),
1142 ];
1143 let hits = relevance_hits("handleList", &docs);
1144 let hit = hits.iter().find(|h| h.id == "hit").unwrap();
1145 assert!(hit.exact_anchor);
1146 assert!(!hits.iter().find(|h| h.id == "noise").unwrap().exact_anchor);
1147 assert_eq!(hits[0].id, "hit");
1149 assert_eq!(hits[0].shape, QueryShape::Identifier);
1150 }
1151
1152 #[test]
1153 fn ranking_arms_are_the_single_list() {
1155 let ids = ranking_arm_ids();
1156 assert!(ids.contains(&"production-blended"));
1157 assert!(ids.contains(&"lexical-then-graph"));
1158 assert!(ids.contains(&"query-routed"));
1159 let mut seen = std::collections::BTreeSet::new();
1160 for id in &ids {
1161 assert!(seen.insert(*id), "duplicate arm {id}");
1162 }
1163 }
1164
1165 #[test]
1166 fn query_mentions_are_precision_first() {
1168 assert!(extract_query_mentions("how does login work").is_empty());
1169 let path = extract_query_mentions("fix sklearn/ensemble/_iforest.py please");
1170 assert_eq!(path.len(), 1);
1171 assert!(path[0].is_path);
1172 assert_eq!(
1173 path[0].segments,
1174 vec!["sklearn", "ensemble", "_iforest.py"]
1175 );
1176 let dotted = extract_query_mentions("see transformers.optimization");
1177 assert_eq!(dotted.len(), 1);
1178 assert!(!dotted[0].is_path);
1179 let tick = extract_query_mentions("look at `handleList` next");
1180 assert_eq!(tick.len(), 1);
1181 assert!(tick[0].backticked);
1182 assert_eq!(tick[0].segments, vec!["handleList"]);
1183 assert!(extract_query_mentions("version 3.10 of python").is_empty());
1184 assert!(extract_query_mentions("e.g. this").is_empty());
1185 }
1186
1187 #[test]
1188 fn query_mentions_mark_anchors_without_changing_bm25() {
1190 let docs = vec![
1191 LexDoc::from_parts("noise", "loginHelper", "src/login.ts", "handles the list", ""),
1192 LexDoc::from_parts("hit", "handleList", "src/server.ts", "", ""),
1193 LexDoc::from_parts("file", "iforest", "sklearn/ensemble/_iforest.py", "", ""),
1194 ];
1195 let hits = relevance_hits("please inspect `handleList` in the service", &docs);
1196 let hit = hits.iter().find(|h| h.id == "hit").unwrap();
1197 assert!(hit.exact_anchor);
1198 let noise = hits.iter().find(|h| h.id == "noise").unwrap();
1199 assert!(!noise.exact_anchor);
1200 let prose = relevance_hits("how does login work", &docs);
1201 assert!(!prose.iter().any(|h| h.exact_anchor));
1202 let path_hits = relevance_hits("bug in sklearn/ensemble/_iforest.py", &docs);
1203 assert!(
1204 path_hits.iter().find(|h| h.id == "file").unwrap().exact_anchor,
1205 "path mention must anchor the named file"
1206 );
1207 }
1208
1209 #[test]
1210 fn persisted_corpus_stats_match_cold_bm25_on_same_docs() {
1212 let docs = vec![
1213 LexDoc::from_parts("a", "handleList", "src/a.py", "list handler", "def handle_list"),
1214 LexDoc::from_parts("b", "other", "src/b.py", "unrelated", "def other"),
1215 LexDoc::from_parts("c", "handleThing", "src/c.py", "", "def handle_thing"),
1216 ];
1217 let stats = Bm25CorpusStats::from_docs(&docs);
1218 assert_eq!(stats.n, 3);
1219 assert!(stats.df.contains_key("handle"));
1220 let cold = bm25_scores("handleList", &docs);
1221 let warm = bm25_scores_with_stats("handleList", &docs, &stats);
1222 assert_eq!(cold.len(), warm.len());
1223 for (c, w) in cold.iter().zip(warm.iter()) {
1224 assert!((c - w).abs() < 1e-9, "cold={c} warm={w}");
1225 }
1226 let cold_hits = relevance_hits("handleList", &docs);
1227 let warm_hits = relevance_hits_with_stats("handleList", &docs, &stats);
1228 assert_eq!(cold_hits.len(), warm_hits.len());
1229 for (c, w) in cold_hits.iter().zip(warm_hits.iter()) {
1230 assert_eq!(c.id, w.id);
1231 assert_eq!(c.exact_anchor, w.exact_anchor);
1232 assert!((c.bm25 - w.bm25).abs() < 1e-9);
1233 }
1234 let subset = vec![docs[0].clone()];
1235 let subset_cold = bm25_scores("handleList", &subset);
1236 let subset_warm = bm25_scores_with_stats("handleList", &subset, &stats);
1237 assert!(
1238 (subset_cold[0] - subset_warm[0]).abs() > 1e-12,
1239 "full-corpus IDF must differ from 1-document slice IDF"
1240 );
1241 }
1242
1243 #[test]
1244 fn path_suffix_match_does_not_hit_extra_py() {
1246 assert!(path_matches_locus("src/app.py", "app.py"));
1247 assert!(path_matches_locus("src/app.py", "src/app.py"));
1248 assert!(!path_matches_locus("src/extra.py", "a.py"));
1249 assert!(!path_matches_locus("src/ba.py", "a.py"));
1250 }
1251}