use std::collections::HashMap;
use std::sync::LazyLock;
fn stem(word: &str) -> String {
use rust_stemmers::{Algorithm, Stemmer};
let stemmer = Stemmer::create(Algorithm::English);
stemmer.stem(word).into_owned()
}
pub fn split_identifier(token: &str) -> Option<Vec<String>> {
if !token.chars().any(char::is_alphanumeric) {
return None;
}
let mut pieces: Vec<&str> = token
.split(|c: char| !c.is_alphanumeric())
.filter(|s| !s.is_empty())
.collect();
let mut subwords: Vec<String> = Vec::new();
if pieces.len() == 1 {
subwords.extend(split_camel_case(pieces.remove(0)));
} else {
for piece in pieces {
subwords.extend(split_camel_case(piece));
}
}
if subwords.len() < 2 {
return None;
}
Some(subwords)
}
fn split_camel_case(run: &str) -> Vec<String> {
if run.is_empty() {
return Vec::new();
}
let chars: Vec<char> = run.chars().collect();
let mut boundaries: Vec<usize> = vec![0];
for i in 1..chars.len() {
let prev = chars[i - 1];
let curr = chars[i];
let lower_to_upper = (prev.is_lowercase() || prev.is_ascii_digit()) && curr.is_uppercase();
let acronym_boundary = prev.is_uppercase()
&& curr.is_uppercase()
&& i + 1 < chars.len()
&& chars[i + 1].is_lowercase();
let alpha_digit = (prev.is_alphabetic() && curr.is_ascii_digit())
|| (prev.is_ascii_digit() && curr.is_alphabetic());
if lower_to_upper || acronym_boundary || alpha_digit {
boundaries.push(i);
}
}
boundaries.push(chars.len());
let mut out = Vec::with_capacity(boundaries.len().saturating_sub(1));
for window in boundaries.windows(2) {
let piece: String = chars[window[0]..window[1]]
.iter()
.collect::<String>()
.to_lowercase();
if !piece.is_empty() {
out.push(piece);
}
}
out
}
fn build_stemmed_synonyms(
raw: &'static [(&'static str, &'static [&'static str])],
) -> HashMap<String, Vec<&'static str>> {
let mut m: HashMap<String, Vec<&'static str>> = HashMap::new();
for &(key, values) in raw {
let stemmed_key: String = key
.split_whitespace()
.map(stem)
.collect::<Vec<_>>()
.join(" ");
m.insert(stemmed_key, values.to_vec());
}
m
}
static SYNONYMS: LazyLock<HashMap<String, Vec<&'static str>>> =
LazyLock::new(|| build_stemmed_synonyms(RAW_SYNONYMS));
const RAW_SYNONYMS: &[(&str, &[&str])] = &[
(
"parallel",
&[
"concurrent",
"async",
"Promise.all",
"allSettled",
"tokio::spawn",
"rayon",
"par_iter",
"ThreadPool",
"ExecutorService",
],
),
(
"concurrent",
&["parallel", "async", "thread", "mutex", "lock", "atomic"],
),
(
"async",
&["await", "Future", "Promise", "Task", "coroutine"],
),
(
"failure",
&["error", "Error", "Result::Err", "reject", "anyhow"],
),
(
"parallel failure",
&["Promise.allSettled", "allSettled", "Promise.all", "settled"],
),
(
"error handling",
&[
"try",
"catch",
"Result",
"Option",
"unwrap_or",
"map_err",
"context",
"bail",
],
),
(
"recovery",
&[
"retry",
"restart",
"resume",
"restore",
"failover",
"resilience",
"fallback",
],
),
(
"fallback",
&["default", "backup", "alternative", "rescue", "degraded"],
),
(
"retry",
&[
"backoff",
"exponential",
"retryWithBackoff",
"max_retries",
"attempt",
"retry_policy",
],
),
(
"connection",
&["connect", "disconnect", "reconnect", "pool", "WebSocket"],
),
(
"lifecycle",
&[
"shutdown",
"dispose",
"cleanup",
"close",
"destroy",
"initiate",
"establish",
"teardown",
"open",
"refresh",
"sync",
],
),
(
"connection lifecycle",
&[
"connect",
"disconnect",
"reconnect",
"connection_pool",
"keep_alive",
"heartbeat",
"idle_timeout",
],
),
(
"pii",
&[
"sanitize",
"redact",
"mask",
"scrub",
"anonymize",
"sensitive",
],
),
(
"redact",
&[
"sanitize",
"mask",
"scrub",
"anonymize",
"censor",
"obfuscate",
"strip",
],
),
(
"authentication",
&[
"auth",
"login",
"logout",
"session",
"token",
"JWT",
"OAuth",
"passport",
"credential",
],
),
(
"authorization",
&[
"permission",
"role",
"access_control",
"ACL",
"RBAC",
"guard",
"policy",
],
),
(
"validation",
&[
"validate",
"schema",
"check",
"verify",
"sanitize",
"constraint",
"Zod",
"Joi",
"pydantic",
],
),
(
"serialization",
&[
"serialize",
"deserialize",
"marshal",
"unmarshal",
"encode",
"decode",
"serde",
"JSON.parse",
"JSON.stringify",
],
),
(
"caching",
&[
"cache",
"memoize",
"memo",
"Redis",
"LRU",
"TTL",
"invalidate",
"cache_control",
],
),
(
"sql",
&[
"SELECT",
"WHERE",
"INSERT",
"QueryBuilder",
"ORM",
"knex",
"prisma",
],
),
(
"query builder",
&["QueryBuilder", "buildQuery", "createQuery", "parameterized"],
),
(
"encryption",
&["encrypt", "decrypt", "cipher", "KMS", "aes", "crypto"],
),
(
"token",
&[
"JWT",
"accessToken",
"refreshToken",
"bearer",
"OAuth",
"credential",
],
),
(
"factory",
&["Factory", "Provider", "builder", "AbstractFactory"],
),
(
"interceptor",
&["middleware", "hook", "filter", "guard", "Interceptor"],
),
("sync", &["synchronize", "replicate", "mirror", "reconcile"]),
(
"middleware",
&[
"interceptor",
"filter",
"guard",
"pipe",
"beforeEach",
"use",
"handler",
],
),
(
"dependency injection",
&[
"inject",
"provider",
"container",
"IoC",
"DI",
"service_locator",
"@Injectable",
],
),
(
"state management",
&[
"store", "reducer", "dispatch", "action", "selector", "Riverpod", "BLoC", "Redux",
"Vuex",
],
),
(
"mock",
&[
"stub",
"fake",
"spy",
"double",
"mockito",
"jest.mock",
"patch",
"monkeypatch",
],
),
(
"test",
&[
"spec",
"assert",
"expect",
"should",
"describe",
"it",
"#[test]",
"def test_",
],
),
(
"thread safety",
&[
"Mutex",
"RwLock",
"atomic",
"Arc",
"synchronized",
"lock",
"guarded",
"thread_local",
],
),
(
"race condition",
&[
"atomic",
"compare_and_swap",
"CAS",
"Mutex",
"ordering",
"fence",
"happens_before",
],
),
(
"memory leak",
&[
"leak", "drop", "Drop", "dispose", "Weak", "finalize", "refcount", "free",
],
),
(
"timeout",
&[
"deadline",
"expires",
"TTL",
"Duration",
"cancel",
"abort",
"WithTimeout",
"set_timeout",
],
),
(
"logging",
&[
"log",
"logger",
"info!",
"debug!",
"warn!",
"tracing",
"println",
"console.log",
"Logger",
"slf4j",
],
),
(
"rate limit",
&[
"throttle",
"debounce",
"quota",
"limiter",
"RateLimiter",
"max_per_second",
"bucket",
],
),
(
"iterator",
&[
"iter",
"Iterator",
"next",
"for_each",
"foreach",
"enumerate",
"Iter",
"yield",
],
),
(
"pagination",
&[
"paginate",
"page",
"offset",
"limit",
"cursor",
"next_page",
"page_size",
"has_more",
],
),
(
"subscription",
&[
"subscribe",
"unsubscribe",
"publisher",
"subscriber",
"observable",
"listen",
"emit",
],
),
(
"event handler",
&[
"on_event",
"handle",
"listener",
"callback",
"addEventListener",
"EventHandler",
"subscribe",
],
),
(
"graceful shutdown",
&[
"shutdown",
"SIGTERM",
"SIGINT",
"cleanup",
"drain",
"close",
"abort_handle",
"cancellation",
],
),
(
"feature flag",
&[
"feature_flag",
"flag",
"toggle",
"rollout",
"experiment",
"if_enabled",
"is_enabled",
],
),
(
"circuit breaker",
&[
"circuit_breaker",
"CircuitBreaker",
"breaker",
"half_open",
"trip",
"failover",
],
),
(
"background job",
&[
"worker",
"job",
"queue",
"scheduler",
"cron",
"tokio::spawn",
"BackgroundService",
"celery",
],
),
(
"telemetry",
&[
"trace",
"span",
"metric",
"OpenTelemetry",
"OTEL",
"Histogram",
"Counter",
"instrument",
],
),
];
pub fn expand_query(query: &str) -> Option<String> {
let raw_words: Vec<&str> = query.split_whitespace().collect();
let query_lower = query.to_lowercase();
let words: Vec<&str> = query_lower.split_whitespace().collect();
let mut id_expansions: Vec<String> = Vec::new();
for raw in &raw_words {
if let Some(subwords) = split_identifier(raw) {
id_expansions.extend(subwords);
}
}
if words.len() < 2 && id_expansions.is_empty() {
return None;
}
let stemmed_words: Vec<String> = words.iter().map(|w| stem(w)).collect();
let mut expansions: Vec<&str> = Vec::new();
let max_window = stemmed_words.len().min(3);
for window_size in (2..=max_window).rev() {
for (i, window) in stemmed_words.windows(window_size).enumerate() {
let phrase = window.join(" ");
if let Some(synonyms) = SYNONYMS.get(&phrase) {
expansions.extend(synonyms.iter());
} else {
let raw_phrase = words[i..i + window_size].join(" ");
if let Some(synonyms) = SYNONYMS.get(&raw_phrase) {
expansions.extend(synonyms.iter());
}
}
}
}
for stemmed in &stemmed_words {
if let Some(synonyms) = SYNONYMS.get(stemmed.as_str()) {
expansions.extend(synonyms.iter());
}
}
if expansions.is_empty() && id_expansions.is_empty() {
return None;
}
expansions.sort_unstable();
expansions.dedup();
id_expansions.sort();
id_expansions.dedup();
let mut tail: Vec<String> = expansions.iter().map(|s| (*s).to_string()).collect();
tail.extend(id_expansions);
Some(format!("{} {}", query, tail.join(" ")))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn expand_parallel_failure_handling() {
let result = expand_query("parallel failure handling").expect("should expand");
assert!(
result.contains("Promise.allSettled"),
"should contain Promise.allSettled"
);
assert!(result.contains("allSettled"), "should contain allSettled");
assert!(result.contains("settled"), "should contain settled");
assert!(result.contains("concurrent"), "should contain concurrent");
assert!(result.contains("error"), "should contain error");
assert!(result.contains("reject"), "should contain reject");
}
#[test]
fn expand_connection_lifecycle() {
let result = expand_query("connection lifecycle").expect("should expand");
assert!(result.contains("connect"), "should contain connect");
assert!(result.contains("disconnect"), "should contain disconnect");
assert!(
result.contains("heartbeat"),
"should contain heartbeat from multi-word match"
);
}
#[test]
fn single_word_returns_none() {
assert!(expand_query("x").is_none());
}
#[test]
fn no_matching_synonyms_returns_none() {
assert!(expand_query("find the bug").is_none());
}
#[test]
fn expansions_are_deduplicated() {
let result = expand_query("parallel concurrent tasks").expect("should expand");
let async_count = result.matches("async").count();
assert_eq!(
async_count, 1,
"async should appear exactly once (deduplicated)"
);
}
#[test]
fn stemmed_lookup_encrypting_matches_encryption() {
let result = expand_query("encrypting sensitive data").expect("should expand");
assert!(
result.contains("KMS"),
"should contain KMS from encryption synonyms"
);
assert!(result.contains("cipher"), "should contain cipher");
assert!(result.contains("decrypt"), "should contain decrypt");
}
#[test]
fn stemmed_lookup_tokens_matches_token() {
let result = expand_query("OAuth tokens refresh").expect("should expand");
assert!(
result.contains("JWT"),
"should contain JWT from token synonyms"
);
assert!(
result.contains("refreshToken"),
"should contain refreshToken"
);
}
#[test]
fn stemmed_lookup_classifying_matches_synonyms() {
let result = expand_query("validating input data");
assert!(
result.is_some(),
"validating should match validation key via stemming"
);
let r = result.expect("should expand");
assert!(
r.contains("sanitize"),
"should contain sanitize from validation synonyms"
);
}
#[test]
fn split_camel_case_basic() {
let parts = split_identifier("getUserById").expect("should split");
assert_eq!(parts, vec!["get", "user", "by", "id"]);
}
#[test]
fn split_pascal_case() {
let parts = split_identifier("RetryPolicy").expect("should split");
assert_eq!(parts, vec!["retry", "policy"]);
}
#[test]
fn split_snake_case() {
let parts = split_identifier("max_retry_count").expect("should split");
assert_eq!(parts, vec!["max", "retry", "count"]);
}
#[test]
fn split_acronym_then_word() {
let parts = split_identifier("HTTPSConnection").expect("should split");
assert_eq!(parts, vec!["https", "connection"]);
}
#[test]
fn split_alpha_digit() {
let parts = split_identifier("OAuth2Provider").expect("should split");
assert!(parts.contains(&"o".to_string()));
assert!(parts.contains(&"auth".to_string()));
assert!(parts.contains(&"2".to_string()));
assert!(parts.contains(&"provider".to_string()));
}
#[test]
fn split_kebab_case() {
let parts = split_identifier("retry-with-backoff").expect("should split");
assert_eq!(parts, vec!["retry", "with", "backoff"]);
}
#[test]
fn split_dot_path() {
let parts = split_identifier("com.example.UserService").expect("should split");
assert_eq!(parts, vec!["com", "example", "user", "service"]);
}
#[test]
fn split_double_colon_path() {
let parts = split_identifier("std::io::Read").expect("should split");
assert_eq!(parts, vec!["std", "io", "read"]);
}
#[test]
fn split_single_lowercase_word_returns_none() {
assert!(split_identifier("user").is_none());
assert!(split_identifier("the").is_none());
}
#[test]
fn split_non_alphanumeric_returns_none() {
assert!(split_identifier("---").is_none());
assert!(split_identifier("").is_none());
}
#[test]
fn expand_query_splits_identifier_in_multiword() {
let result = expand_query("how does getUserById work").expect("should expand");
assert!(result.contains(" get "), "subword 'get' should appear");
assert!(result.contains(" user"), "subword 'user' should appear");
assert!(result.contains(" by "), "subword 'by' should appear");
assert!(result.contains(" id"), "subword 'id' should appear");
}
#[test]
fn expand_query_splits_single_identifier_only() {
assert!(expand_query("hello").is_none());
let result = expand_query("RetryPolicy");
assert!(
result.is_some(),
"single-word identifier should expand via subwords"
);
let r = result.expect("should expand");
assert!(r.contains("retry"));
assert!(r.contains("policy"));
}
#[test]
fn expand_thread_safety_pair() {
let r = expand_query("thread safety concerns").expect("should expand");
assert!(r.contains("Mutex"));
assert!(r.contains("atomic"));
}
#[test]
fn expand_rate_limit_pair() {
let r = expand_query("rate limit api calls").expect("should expand");
assert!(r.contains("throttle"));
assert!(r.contains("RateLimiter"));
}
#[test]
fn expand_graceful_shutdown_pair() {
let r = expand_query("graceful shutdown sequence").expect("should expand");
assert!(r.contains("SIGTERM"));
assert!(r.contains("drain"));
}
#[test]
fn expand_memory_leak_pair() {
let r = expand_query("debug memory leak").expect("should expand");
assert!(r.contains("Weak"));
assert!(r.contains("dispose"));
}
#[test]
fn expand_telemetry_pair() {
let r = expand_query("add telemetry to handler").expect("should expand");
assert!(r.contains("trace"));
assert!(r.contains("OpenTelemetry"));
}
#[test]
fn expand_dual_route_differs_from_single_route() {
let original = "thread safety check";
let expanded = expand_query(original).expect("should expand");
assert!(
expanded.starts_with(original),
"expanded should preserve original prefix"
);
assert!(
expanded.len() > original.len(),
"expanded should be strictly longer than original"
);
assert!(expanded.contains("Mutex"));
}
}