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    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}