Skip to main content

harn_session_store/
search.rs

1//! Canonical session/transcript search contract.
2//!
3//! Ranking, scope, fallback reporting, and searchable-text projection live
4//! here so storage adapters and transports cannot grow competing policy.
5
6use 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
104/// Backend-neutral embedding seam used by the store's search implementation.
105///
106/// The deterministic lexical backend is always available. Higher-quality
107/// implementations may be injected through `StoreHooks` without changing the
108/// search interface or any transport.
109pub 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
122/// Deterministic cross-platform lexical-hash floor.
123pub 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}