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 #[expect(
296 clippy::string_slice,
297 reason = "original_byte_anchor is walked back to a char boundary by the loop above"
298 )]
299 let char_anchor = text[..original_byte_anchor].chars().count();
300 let start = char_anchor.saturating_sub(max_chars / 3);
301 let excerpt = text.chars().skip(start).take(max_chars).collect::<String>();
302 format!(
303 "{}{}{}",
304 if start > 0 { "…" } else { "" },
305 excerpt,
306 if start + max_chars < text.chars().count() {
307 "…"
308 } else {
309 ""
310 }
311 )
312}
313
314pub(crate) fn lexical_score(query: &str, text: &str) -> f32 {
315 let query_tokens = word_tokens(query);
316 if query_tokens.is_empty() {
317 return 0.0;
318 }
319 let text_tokens = word_tokens(text);
320 let frequencies =
321 text_tokens
322 .into_iter()
323 .fold(BTreeMap::<String, usize>::new(), |mut counts, token| {
324 *counts.entry(token).or_default() += 1;
325 counts
326 });
327 if query_tokens
328 .iter()
329 .any(|token| !frequencies.contains_key(token))
330 {
331 return 0.0;
332 }
333 let matched = query_tokens
334 .iter()
335 .filter_map(|token| frequencies.get(token))
336 .map(|count| 1.0 + (*count as f32).ln())
337 .sum::<f32>();
338 let exact = text
339 .to_lowercase()
340 .contains(query.trim().to_lowercase().as_str());
341 matched / query_tokens.len() as f32 + if exact { 1.0 } else { 0.0 }
342}
343
344pub(crate) fn combined_score(
345 mode: SearchMode,
346 fts_rank: Option<usize>,
347 semantic_rank: Option<usize>,
348 fts_score: Option<f32>,
349 semantic_score: Option<f32>,
350) -> f32 {
351 match mode {
352 SearchMode::Fts => fts_score.unwrap_or_default(),
353 SearchMode::Semantic => semantic_score.unwrap_or_default(),
354 SearchMode::Hybrid => {
355 fts_rank
356 .map(|rank| 1.0 / (RRF_K + rank as f32 + 1.0))
357 .unwrap_or_default()
358 + semantic_rank
359 .map(|rank| 1.0 / (RRF_K + rank as f32 + 1.0))
360 .unwrap_or_default()
361 }
362 }
363}
364
365pub(crate) fn ranks(scores: &[f32]) -> BTreeMap<usize, usize> {
366 let mut ranked = scores
367 .iter()
368 .copied()
369 .enumerate()
370 .filter(|(_, score)| *score > 0.0)
371 .collect::<Vec<_>>();
372 ranked.sort_by(|(left_index, left), (right_index, right)| {
373 right
374 .total_cmp(left)
375 .then_with(|| left_index.cmp(right_index))
376 });
377 ranked
378 .into_iter()
379 .enumerate()
380 .map(|(rank, (index, _))| (index, rank))
381 .collect()
382}
383
384pub(crate) fn vector_blob(vector: &[f32]) -> Vec<u8> {
385 let mut bytes = Vec::with_capacity(std::mem::size_of_val(vector));
386 for value in vector {
387 bytes.extend_from_slice(&value.to_le_bytes());
388 }
389 bytes
390}
391
392pub(crate) fn vector_from_blob(bytes: &[u8], dim: usize) -> Option<Vec<f32>> {
393 if bytes.len() != dim.checked_mul(std::mem::size_of::<f32>())? {
394 return None;
395 }
396 Some(
397 bytes
398 .chunks_exact(4)
399 .map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
400 .collect(),
401 )
402}
403
404pub(crate) fn word_tokens(text: &str) -> Vec<String> {
405 let mut tokens = Vec::new();
406 let mut current = String::new();
407 let mut previous_lower = false;
408 let flush = |current: &mut String, tokens: &mut Vec<String>| {
409 if !current.is_empty() {
410 tokens.push(std::mem::take(current));
411 }
412 };
413 for character in text.chars() {
414 if character.is_alphanumeric() {
415 if character.is_uppercase() && previous_lower {
416 flush(&mut current, &mut tokens);
417 }
418 current.extend(character.to_lowercase());
419 previous_lower = character.is_lowercase() || character.is_numeric();
420 } else {
421 flush(&mut current, &mut tokens);
422 previous_lower = false;
423 }
424 }
425 flush(&mut current, &mut tokens);
426 tokens
427}
428
429pub(crate) fn fts_literal_query(query: &str) -> String {
430 word_tokens(query)
431 .into_iter()
432 .map(|token| format!("\"{}\"", token.replace('"', "\"\"")))
433 .collect::<Vec<_>>()
434 .join(" AND ")
435}
436
437fn char_ngrams(text: &str, width: usize) -> Vec<String> {
438 if width == 0 {
439 return Vec::new();
440 }
441 let mut normalized = String::with_capacity(text.len() + 2);
442 normalized.push(' ');
443 let mut previous_space = true;
444 for character in text.chars() {
445 if character.is_whitespace() {
446 if !previous_space {
447 normalized.push(' ');
448 previous_space = true;
449 }
450 } else {
451 normalized.extend(character.to_lowercase());
452 previous_space = false;
453 }
454 }
455 if !previous_space {
456 normalized.push(' ');
457 }
458 let characters = normalized.chars().collect::<Vec<_>>();
459 characters
460 .windows(width)
461 .map(|window| window.iter().collect())
462 .collect()
463}
464
465fn fnv1a(bytes: &[u8], seed: u64) -> u64 {
466 const FNV_PRIME: u64 = 0x0000_0100_0000_01B3;
467 let mut hash = seed ^ 0xcbf2_9ce4_8422_2325;
468 for byte in bytes {
469 hash ^= u64::from(*byte);
470 hash = hash.wrapping_mul(FNV_PRIME);
471 }
472 hash
473}
474
475fn collect_json_strings(value: &serde_json::Value, parts: &mut Vec<String>) {
476 match value {
477 serde_json::Value::String(text) => parts.push(text.clone()),
478 serde_json::Value::Array(items) => {
479 for item in items {
480 collect_json_strings(item, parts);
481 }
482 }
483 serde_json::Value::Object(fields) => {
484 for value in fields.values() {
485 collect_json_strings(value, parts);
486 }
487 }
488 serde_json::Value::Null | serde_json::Value::Bool(_) | serde_json::Value::Number(_) => {}
489 }
490}
491
492#[cfg(test)]
493mod tests {
494 use super::*;
495
496 #[test]
497 fn lexical_embedder_is_deterministic_and_related() {
498 let embedder = LexicalEmbedder::default();
499 let query = embedder.embed("rate limiting middleware");
500 assert_eq!(query, embedder.embed("rate limiting middleware"));
501 assert!(
502 cosine(&query, &embedder.embed("API rate limiter"))
503 > cosine(&query, &embedder.embed("markdown table renderer"))
504 );
505 }
506
507 #[test]
508 fn fts_queries_are_literal_and_identifier_aware() {
509 assert_eq!(
510 fts_literal_query("getUserByID OR token*"),
511 "\"get\" AND \"user\" AND \"by\" AND \"id\" AND \"or\" AND \"token\""
512 );
513 }
514
515 #[test]
516 fn vector_blob_round_trips() {
517 let vector = vec![-1.0, 0.25, 4.0];
518 assert_eq!(vector_from_blob(&vector_blob(&vector), 3), Some(vector));
519 assert_eq!(vector_from_blob(&[0, 1], 3), None);
520 }
521
522 #[test]
523 fn search_requires_an_explicit_scope() {
524 let error = SearchQuery {
525 query: "needle".to_string(),
526 mode: SearchMode::Fts,
527 filter: SearchFilter::default(),
528 limit: None,
529 }
530 .validate()
531 .expect_err("unscoped search must be rejected");
532 assert!(error.contains("requires"));
533 }
534
535 #[test]
536 fn unicode_snippet_anchor_never_slices_at_a_folded_byte_offset() {
537 let text = format!("{}needle", "İ".repeat(300));
538 let rendered = snippet(&text, "needle", 40);
539 assert!(rendered.contains("needle"));
540 }
541}