1use std::collections::BTreeMap;
7use std::sync::Arc;
8
9use serde::{Deserialize, Serialize};
10
11use crate::redaction::SharedEventRedactor;
12use crate::{EventId, SessionEventKind, SessionMeta, StoredEvent};
13
14pub const DEFAULT_SEARCH_LIMIT: usize = 50;
15pub const MAX_SEARCH_LIMIT: usize = 500;
16const RRF_K: f32 = 60.0;
17
18#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
19#[serde(rename_all = "snake_case")]
20pub enum SearchMode {
21 Fts,
22 Semantic,
23 #[default]
24 Hybrid,
25}
26
27#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
28pub struct SearchFilter {
29 #[serde(default)]
30 pub tenant_id: Option<String>,
31 #[serde(default)]
32 pub project_scope: Option<String>,
33 #[serde(default)]
34 pub session_id: Option<String>,
35}
36
37#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
38pub struct SearchQuery {
39 pub query: String,
40 #[serde(default)]
41 pub mode: SearchMode,
42 #[serde(default)]
43 pub filter: SearchFilter,
44 #[serde(default)]
45 pub limit: Option<usize>,
46}
47
48impl SearchQuery {
49 pub fn validate(&self) -> Result<(), String> {
50 if self.query.trim().is_empty() {
51 return Err("search query must be non-empty".to_string());
52 }
53 if self.query.chars().any(|character| character == '\0') {
54 return Err("search query must not contain NUL".to_string());
55 }
56 let has_scope = [
57 self.filter.tenant_id.as_deref(),
58 self.filter.project_scope.as_deref(),
59 self.filter.session_id.as_deref(),
60 ]
61 .into_iter()
62 .flatten()
63 .any(|scope| !scope.trim().is_empty());
64 if !has_scope {
65 return Err(
66 "search requires tenant_id, project_scope, or session_id scope".to_string(),
67 );
68 }
69 Ok(())
70 }
71
72 pub fn limit(&self) -> usize {
73 self.limit
74 .unwrap_or(DEFAULT_SEARCH_LIMIT)
75 .clamp(1, MAX_SEARCH_LIMIT)
76 }
77}
78
79#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
80pub struct SearchHit {
81 pub session_id: String,
82 pub event_id: EventId,
83 pub kind: SessionEventKind,
84 pub score: f32,
85 #[serde(skip_serializing_if = "Option::is_none")]
86 pub fts_score: Option<f32>,
87 #[serde(skip_serializing_if = "Option::is_none")]
88 pub semantic_score: Option<f32>,
89 pub snippet: String,
90 pub event: StoredEvent,
91}
92
93#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
94pub struct SearchResponse {
95 pub requested_mode: SearchMode,
96 pub effective_mode: SearchMode,
97 pub embedding_backend: String,
98 pub semantic_floor: bool,
99 #[serde(skip_serializing_if = "Option::is_none")]
100 pub fallback_reason: Option<String>,
101 pub hits: Vec<SearchHit>,
102}
103
104pub trait Embedder: Send + Sync {
110 fn embed(&self, text: &str) -> Vec<f32>;
111 fn dim(&self) -> usize;
112 fn name(&self) -> &str;
113 fn is_semantic(&self) -> bool {
114 true
115 }
116
117 fn embed_batch(&self, texts: &[String]) -> Vec<Vec<f32>> {
118 texts.iter().map(|text| self.embed(text)).collect()
119 }
120}
121
122pub struct LexicalEmbedder {
124 dim: usize,
125}
126
127impl LexicalEmbedder {
128 pub fn new(dim: usize) -> Self {
129 Self { dim: dim.max(16) }
130 }
131
132 fn add_feature(&self, vector: &mut [f32], feature: &str, weight: f32) {
133 let hash = fnv1a(feature.as_bytes(), 0);
134 let bucket = (hash % self.dim as u64) as usize;
135 let sign = if fnv1a(feature.as_bytes(), 0x9e37_79b9_7f4a_7c15) & 1 == 0 {
136 1.0
137 } else {
138 -1.0
139 };
140 vector[bucket] += sign * weight;
141 }
142}
143
144impl Default for LexicalEmbedder {
145 fn default() -> Self {
146 Self::new(256)
147 }
148}
149
150impl Embedder for LexicalEmbedder {
151 fn embed(&self, text: &str) -> Vec<f32> {
152 let mut vector = vec![0.0; self.dim];
153 for token in word_tokens(text) {
154 self.add_feature(&mut vector, &token, 1.0);
155 }
156 for gram in char_ngrams(text, 3) {
157 self.add_feature(&mut vector, &gram, 0.35);
158 }
159 l2_normalize(&mut vector);
160 vector
161 }
162
163 fn dim(&self) -> usize {
164 self.dim
165 }
166
167 #[allow(clippy::unnecessary_literal_bound)]
168 fn name(&self) -> &str {
169 "lexical-hash"
170 }
171
172 fn is_semantic(&self) -> bool {
173 false
174 }
175}
176
177pub fn default_embedder() -> Arc<dyn Embedder> {
178 Arc::new(LexicalEmbedder::default())
179}
180
181pub fn cosine(left: &[f32], right: &[f32]) -> f32 {
182 if left.is_empty() || left.len() != right.len() {
183 return 0.0;
184 }
185 let mut dot = 0.0;
186 let mut left_norm = 0.0;
187 let mut right_norm = 0.0;
188 for (left, right) in left.iter().zip(right.iter()) {
189 if !left.is_finite() || !right.is_finite() {
190 return 0.0;
191 }
192 dot += left * right;
193 left_norm += left * left;
194 right_norm += right * right;
195 }
196 if left_norm <= 0.0 || right_norm <= 0.0 {
197 return 0.0;
198 }
199 (dot / (left_norm.sqrt() * right_norm.sqrt())).clamp(-1.0, 1.0)
200}
201
202pub fn l2_normalize(vector: &mut [f32]) {
203 let norm = vector.iter().map(|value| value * value).sum::<f32>().sqrt();
204 if norm > 0.0 {
205 for value in vector {
206 *value /= norm;
207 }
208 }
209}
210
211pub fn event_search_text(event: &StoredEvent) -> String {
212 let mut parts = Vec::new();
213 parts.push(event.kind.discriminator().replace('_', " "));
214 if let Some(actor) = event.actor.as_deref() {
215 parts.push(actor.to_string());
216 }
217 collect_json_strings(&event.payload, &mut parts);
218 parts.join("\n")
219}
220
221pub(crate) fn redacted_search_document(
222 redactor: Option<&SharedEventRedactor>,
223 meta: &SessionMeta,
224 event: &StoredEvent,
225) -> String {
226 redacted_search_document_parts(
227 redactor,
228 meta.title.as_deref(),
229 meta.cwd.as_deref(),
230 meta.model.as_deref(),
231 meta.project_scope.as_deref(),
232 event,
233 )
234}
235
236pub(crate) fn redacted_search_document_parts(
237 redactor: Option<&SharedEventRedactor>,
238 title: Option<&str>,
239 cwd: Option<&str>,
240 model: Option<&str>,
241 project_scope: Option<&str>,
242 event: &StoredEvent,
243) -> String {
244 let mut metadata = serde_json::json!({
245 "title": title,
246 "cwd": cwd,
247 "model": model,
248 "project_scope": project_scope,
249 });
250 if let Some(redactor) = redactor {
251 redactor.redact_json_in_place(&mut metadata);
252 }
253 search_document_parts(
254 metadata.get("title").and_then(serde_json::Value::as_str),
255 metadata.get("cwd").and_then(serde_json::Value::as_str),
256 metadata.get("model").and_then(serde_json::Value::as_str),
257 metadata
258 .get("project_scope")
259 .and_then(serde_json::Value::as_str),
260 event,
261 )
262}
263
264pub(crate) fn search_document_parts(
265 title: Option<&str>,
266 cwd: Option<&str>,
267 model: Option<&str>,
268 project_scope: Option<&str>,
269 event: &StoredEvent,
270) -> String {
271 let event_text = event_search_text(event);
272 [title, cwd, model, project_scope, Some(event_text.as_str())]
273 .into_iter()
274 .flatten()
275 .collect::<Vec<_>>()
276 .join("\n")
277}
278
279pub fn snippet(text: &str, query: &str, max_chars: usize) -> String {
280 let text = text.trim();
281 if text.chars().count() <= max_chars {
282 return text.to_string();
283 }
284 let folded = text.to_lowercase();
285 let needle = word_tokens(query).into_iter().next().unwrap_or_default();
286 let byte_anchor = if needle.is_empty() {
287 0
288 } else {
289 folded.find(&needle).unwrap_or(0)
290 };
291 let mut original_byte_anchor = byte_anchor.min(text.len());
292 while original_byte_anchor > 0 && !text.is_char_boundary(original_byte_anchor) {
293 original_byte_anchor -= 1;
294 }
295 let char_anchor = text[..original_byte_anchor].chars().count();
296 let start = char_anchor.saturating_sub(max_chars / 3);
297 let excerpt = text.chars().skip(start).take(max_chars).collect::<String>();
298 format!(
299 "{}{}{}",
300 if start > 0 { "…" } else { "" },
301 excerpt,
302 if start + max_chars < text.chars().count() {
303 "…"
304 } else {
305 ""
306 }
307 )
308}
309
310pub(crate) fn lexical_score(query: &str, text: &str) -> f32 {
311 let query_tokens = word_tokens(query);
312 if query_tokens.is_empty() {
313 return 0.0;
314 }
315 let text_tokens = word_tokens(text);
316 let frequencies =
317 text_tokens
318 .into_iter()
319 .fold(BTreeMap::<String, usize>::new(), |mut counts, token| {
320 *counts.entry(token).or_default() += 1;
321 counts
322 });
323 if query_tokens
324 .iter()
325 .any(|token| !frequencies.contains_key(token))
326 {
327 return 0.0;
328 }
329 let matched = query_tokens
330 .iter()
331 .filter_map(|token| frequencies.get(token))
332 .map(|count| 1.0 + (*count as f32).ln())
333 .sum::<f32>();
334 let exact = text
335 .to_lowercase()
336 .contains(query.trim().to_lowercase().as_str());
337 matched / query_tokens.len() as f32 + if exact { 1.0 } else { 0.0 }
338}
339
340pub(crate) fn combined_score(
341 mode: SearchMode,
342 fts_rank: Option<usize>,
343 semantic_rank: Option<usize>,
344 fts_score: Option<f32>,
345 semantic_score: Option<f32>,
346) -> f32 {
347 match mode {
348 SearchMode::Fts => fts_score.unwrap_or_default(),
349 SearchMode::Semantic => semantic_score.unwrap_or_default(),
350 SearchMode::Hybrid => {
351 fts_rank
352 .map(|rank| 1.0 / (RRF_K + rank as f32 + 1.0))
353 .unwrap_or_default()
354 + semantic_rank
355 .map(|rank| 1.0 / (RRF_K + rank as f32 + 1.0))
356 .unwrap_or_default()
357 }
358 }
359}
360
361pub(crate) fn ranks(scores: &[f32]) -> BTreeMap<usize, usize> {
362 let mut ranked = scores
363 .iter()
364 .copied()
365 .enumerate()
366 .filter(|(_, score)| *score > 0.0)
367 .collect::<Vec<_>>();
368 ranked.sort_by(|(left_index, left), (right_index, right)| {
369 right
370 .total_cmp(left)
371 .then_with(|| left_index.cmp(right_index))
372 });
373 ranked
374 .into_iter()
375 .enumerate()
376 .map(|(rank, (index, _))| (index, rank))
377 .collect()
378}
379
380pub(crate) fn vector_blob(vector: &[f32]) -> Vec<u8> {
381 let mut bytes = Vec::with_capacity(std::mem::size_of_val(vector));
382 for value in vector {
383 bytes.extend_from_slice(&value.to_le_bytes());
384 }
385 bytes
386}
387
388pub(crate) fn vector_from_blob(bytes: &[u8], dim: usize) -> Option<Vec<f32>> {
389 if bytes.len() != dim.checked_mul(std::mem::size_of::<f32>())? {
390 return None;
391 }
392 Some(
393 bytes
394 .chunks_exact(4)
395 .map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
396 .collect(),
397 )
398}
399
400pub(crate) fn word_tokens(text: &str) -> Vec<String> {
401 let mut tokens = Vec::new();
402 let mut current = String::new();
403 let mut previous_lower = false;
404 let flush = |current: &mut String, tokens: &mut Vec<String>| {
405 if !current.is_empty() {
406 tokens.push(std::mem::take(current));
407 }
408 };
409 for character in text.chars() {
410 if character.is_alphanumeric() {
411 if character.is_uppercase() && previous_lower {
412 flush(&mut current, &mut tokens);
413 }
414 current.extend(character.to_lowercase());
415 previous_lower = character.is_lowercase() || character.is_numeric();
416 } else {
417 flush(&mut current, &mut tokens);
418 previous_lower = false;
419 }
420 }
421 flush(&mut current, &mut tokens);
422 tokens
423}
424
425pub(crate) fn fts_literal_query(query: &str) -> String {
426 word_tokens(query)
427 .into_iter()
428 .map(|token| format!("\"{}\"", token.replace('"', "\"\"")))
429 .collect::<Vec<_>>()
430 .join(" AND ")
431}
432
433fn char_ngrams(text: &str, width: usize) -> Vec<String> {
434 if width == 0 {
435 return Vec::new();
436 }
437 let mut normalized = String::with_capacity(text.len() + 2);
438 normalized.push(' ');
439 let mut previous_space = true;
440 for character in text.chars() {
441 if character.is_whitespace() {
442 if !previous_space {
443 normalized.push(' ');
444 previous_space = true;
445 }
446 } else {
447 normalized.extend(character.to_lowercase());
448 previous_space = false;
449 }
450 }
451 if !previous_space {
452 normalized.push(' ');
453 }
454 let characters = normalized.chars().collect::<Vec<_>>();
455 characters
456 .windows(width)
457 .map(|window| window.iter().collect())
458 .collect()
459}
460
461fn fnv1a(bytes: &[u8], seed: u64) -> u64 {
462 const FNV_PRIME: u64 = 0x0000_0100_0000_01B3;
463 let mut hash = seed ^ 0xcbf2_9ce4_8422_2325;
464 for byte in bytes {
465 hash ^= u64::from(*byte);
466 hash = hash.wrapping_mul(FNV_PRIME);
467 }
468 hash
469}
470
471fn collect_json_strings(value: &serde_json::Value, parts: &mut Vec<String>) {
472 match value {
473 serde_json::Value::String(text) => parts.push(text.clone()),
474 serde_json::Value::Array(items) => {
475 for item in items {
476 collect_json_strings(item, parts);
477 }
478 }
479 serde_json::Value::Object(fields) => {
480 for value in fields.values() {
481 collect_json_strings(value, parts);
482 }
483 }
484 serde_json::Value::Null | serde_json::Value::Bool(_) | serde_json::Value::Number(_) => {}
485 }
486}
487
488#[cfg(test)]
489mod tests {
490 use super::*;
491
492 #[test]
493 fn lexical_embedder_is_deterministic_and_related() {
494 let embedder = LexicalEmbedder::default();
495 let query = embedder.embed("rate limiting middleware");
496 assert_eq!(query, embedder.embed("rate limiting middleware"));
497 assert!(
498 cosine(&query, &embedder.embed("API rate limiter"))
499 > cosine(&query, &embedder.embed("markdown table renderer"))
500 );
501 }
502
503 #[test]
504 fn fts_queries_are_literal_and_identifier_aware() {
505 assert_eq!(
506 fts_literal_query("getUserByID OR token*"),
507 "\"get\" AND \"user\" AND \"by\" AND \"id\" AND \"or\" AND \"token\""
508 );
509 }
510
511 #[test]
512 fn vector_blob_round_trips() {
513 let vector = vec![-1.0, 0.25, 4.0];
514 assert_eq!(vector_from_blob(&vector_blob(&vector), 3), Some(vector));
515 assert_eq!(vector_from_blob(&[0, 1], 3), None);
516 }
517
518 #[test]
519 fn search_requires_an_explicit_scope() {
520 let error = SearchQuery {
521 query: "needle".to_string(),
522 mode: SearchMode::Fts,
523 filter: SearchFilter::default(),
524 limit: None,
525 }
526 .validate()
527 .expect_err("unscoped search must be rejected");
528 assert!(error.contains("requires"));
529 }
530
531 #[test]
532 fn unicode_snippet_anchor_never_slices_at_a_folded_byte_offset() {
533 let text = format!("{}needle", "İ".repeat(300));
534 let rendered = snippet(&text, "needle", 40);
535 assert!(rendered.contains("needle"));
536 }
537}