Skip to main content

omgbase_search/
embed.rs

1//! Embeddings (`spec/search` §2): which blocks embed, the exact strings a
2//! provider is asked for, the cache keys, the document input and its
3//! whole-vs-pooled budget, and pooling.
4
5use sha2::{Digest, Sha256};
6
7use crate::fts::split_ws;
8
9/// §2.1: the minimum token count for a block to embed on its own.
10pub const MIN_EMBED_TOKENS: usize = 24;
11/// §2.3: cache misses are embedded in requests of this many inputs.
12pub const EMBED_BATCH: usize = 32;
13/// §2.4: the whole-document budget when the provider reports no limit.
14pub const DEFAULT_DOC_TOKEN_BUDGET: u32 = 512;
15/// §2.4: tokens reserved for the document header line.
16pub const DOC_HEADER_MARGIN_TOKENS: u32 = 16;
17
18/// `sha256(utf8(s))`.
19#[must_use]
20pub fn sha256(s: &str) -> [u8; 32] {
21    Sha256::digest(s.as_bytes()).into()
22}
23
24/// Lower-case hex of `bytes`.
25#[must_use]
26pub fn hex(bytes: &[u8]) -> String {
27    bytes.iter().map(|b| format!("{b:02x}")).collect()
28}
29
30/// §2.1: the number of non-empty pieces of `text` split on `\s+`.
31#[must_use]
32pub fn word_count(text: &str) -> usize {
33    split_ws(text).count()
34}
35
36/// §2.1: a block embeds on its own when its text has at least
37/// [`MIN_EMBED_TOKENS`] words.
38#[must_use]
39pub fn should_embed(text: &str) -> bool {
40    word_count(text) >= MIN_EMBED_TOKENS
41}
42
43/// §2.1: `ceil(words × 1.3)`, the budget estimate and the pooling weight.
44#[must_use]
45pub fn estimate_tokens(text: &str) -> u64 {
46    (word_count(text) as f64 * 1.3).ceil() as u64
47}
48
49/// §2.2: `doc_title · path · chain.join(" › ") · type`.
50#[must_use]
51pub fn context_prefix(
52    doc_title: &str,
53    path: &str,
54    heading_chain: &[String],
55    block_type: &str,
56) -> String {
57    format!(
58        "{doc_title} \u{00B7} {path} \u{00B7} {} \u{00B7} {block_type}",
59        heading_chain.join(" \u{203A} ")
60    )
61}
62
63/// §2.2: what the provider is asked for: `ctx + "\n" + text`.
64#[must_use]
65pub fn embed_input(ctx: &str, block_text: &str) -> String {
66    format!("{ctx}\n{block_text}")
67}
68
69/// §2.2: `sha256(ctx)`, the second cache key.
70#[must_use]
71pub fn ctx_hash(ctx: &str) -> [u8; 32] {
72    sha256(ctx)
73}
74
75/// One block to embed (§2.2–§2.3).
76#[derive(Clone, Debug, PartialEq, Eq)]
77pub struct EmbedTask {
78    pub block_id: String,
79    /// The block's `raw_hash`, hex.
80    pub content_hash: String,
81    pub ctx: String,
82    pub text: String,
83}
84
85impl EmbedTask {
86    /// The provider input for this task.
87    #[must_use]
88    pub fn input(&self) -> String {
89        embed_input(&self.ctx, &self.text)
90    }
91}
92
93/// §2.4: `[title, path, "type: " + type?, "layer: " + layer?].join(" · ")`;
94/// `type`/`layer` only when given (the caller passes the merged property when
95/// it is a non-blank string).
96#[must_use]
97pub fn doc_header(title: &str, path: &str, doc_type: Option<&str>, layer: Option<&str>) -> String {
98    let mut bits = vec![title.to_owned(), path.to_owned()];
99    if let Some(t) = doc_type {
100        bits.push(format!("type: {t}"));
101    }
102    if let Some(l) = layer {
103        bits.push(format!("layer: {l}"));
104    }
105    bits.join(" \u{00B7} ")
106}
107
108/// §2.4: `header + "\n" + body`.
109#[must_use]
110pub fn doc_input(header: &str, body: &str) -> String {
111    format!("{header}\n{body}")
112}
113
114/// §2.4: `max(1, (max_input_tokens ?? 512) − 16)`.
115#[must_use]
116pub fn token_budget(max_input_tokens: Option<u32>) -> u64 {
117    let budget = max_input_tokens.unwrap_or(DEFAULT_DOC_TOKEN_BUDGET);
118    u64::from(budget.saturating_sub(DOC_HEADER_MARGIN_TOKENS).max(1))
119}
120
121/// How a document vector was (or would be) computed (§2.4).
122#[derive(Clone, Copy, Debug, PartialEq, Eq)]
123pub enum DocEmbedMethod {
124    /// The input was within budget and sent to the provider.
125    Whole,
126    /// Over budget: the token-weighted mean of the cached block vectors.
127    Pooled,
128}
129
130impl DocEmbedMethod {
131    /// The `doc_embeddings.method` spelling.
132    #[must_use]
133    pub fn as_str(self) -> &'static str {
134        match self {
135            DocEmbedMethod::Whole => "whole",
136            DocEmbedMethod::Pooled => "pooled",
137        }
138    }
139
140    /// Parse the stored spelling.
141    #[must_use]
142    pub fn parse(s: &str) -> Option<Self> {
143        match s {
144            "whole" => Some(DocEmbedMethod::Whole),
145            "pooled" => Some(DocEmbedMethod::Pooled),
146            _ => None,
147        }
148    }
149
150    /// §2.4: whole when `estimate_tokens(input)` is within `budget`, else pooled.
151    #[must_use]
152    pub fn for_input(input: &str, budget: u64) -> Self {
153        if estimate_tokens(input) <= budget {
154            DocEmbedMethod::Whole
155        } else {
156            DocEmbedMethod::Pooled
157        }
158    }
159}
160
161/// A block's contribution to a pooled document vector (§2.5): where its cached
162/// vector is keyed and its weight.
163#[derive(Clone, Debug, PartialEq, Eq)]
164pub struct DocEmbedBlockRef {
165    /// The block's `raw_hash`, hex.
166    pub content_hash: String,
167    pub ctx: String,
168    /// `estimate_tokens(text)`.
169    pub tokens: u64,
170}
171
172/// One document to embed (§2.4).
173#[derive(Clone, Debug, PartialEq, Eq)]
174pub struct DocEmbedTask {
175    pub doc_id: String,
176    /// The §2.4 header line.
177    pub header: String,
178    /// `header + "\n" + reconstruct(doc)`.
179    pub input: String,
180    /// The document's embeddable blocks in `(path, ordinal)` order.
181    pub blocks: Vec<DocEmbedBlockRef>,
182}
183
184impl DocEmbedTask {
185    /// §2.4: `sha256(input)`, the freshness key.
186    #[must_use]
187    pub fn input_hash(&self) -> [u8; 32] {
188        sha256(&self.input)
189    }
190}
191
192/// §2.5: over `refs` in order, each with weight `w = max(1, tokens)`, skip
193/// blocks with no cached vector; `acc[i] += w × v[i]` over `min(dim, |v|)`
194/// components in f64; `None` when nothing contributed; else `acc[i] /= Σw`,
195/// `norm = √Σacc²`, `out[i] = acc[i] / norm` as float32 (zeros when the norm
196/// is 0).
197pub fn pool_block_vectors<F>(
198    dim: usize,
199    refs: &[DocEmbedBlockRef],
200    mut cached: F,
201) -> Option<Vec<f32>>
202where
203    F: FnMut(&DocEmbedBlockRef) -> Option<Vec<f32>>,
204{
205    let mut acc = vec![0.0f64; dim];
206    let mut weight_sum = 0.0f64;
207    for r in refs {
208        let Some(v) = cached(r) else {
209            continue;
210        };
211        let w = if r.tokens > 0 { r.tokens as f64 } else { 1.0 };
212        let n = dim.min(v.len());
213        for i in 0..n {
214            acc[i] += w * f64::from(v[i]);
215        }
216        weight_sum += w;
217    }
218    if weight_sum == 0.0 {
219        return None;
220    }
221    let mut norm = 0.0f64;
222    for a in &mut acc {
223        *a /= weight_sum;
224        norm += *a * *a;
225    }
226    let norm = norm.sqrt();
227    if norm == 0.0 {
228        return Some(vec![0.0f32; dim]);
229    }
230    Some(acc.iter().map(|a| (a / norm) as f32).collect())
231}
232
233#[cfg(test)]
234mod tests {
235    use super::*;
236
237    #[test]
238    fn should_embed_threshold_and_estimate() {
239        let words = |n: usize| {
240            (0..n)
241                .map(|i| format!("w{i}"))
242                .collect::<Vec<_>>()
243                .join(" ")
244        };
245        assert!(!should_embed(&words(23)));
246        assert!(should_embed(&words(24)));
247        assert!(should_embed(&format!("  {}  \n", words(24))));
248        assert!(!should_embed(""));
249        assert_eq!(estimate_tokens(""), 0);
250        assert_eq!(estimate_tokens("one"), 2);
251        assert_eq!(estimate_tokens("one two three"), 4); // ceil(3.9)
252        assert_eq!(estimate_tokens(&words(10)), 13);
253        assert_eq!(estimate_tokens(&words(24)), 32); // ceil(31.200000000000003)
254    }
255
256    #[test]
257    fn context_prefix_shape() {
258        assert_eq!(
259            context_prefix(
260                "T",
261                "a.md",
262                &["H1".to_owned(), "H2".to_owned()],
263                "paragraph"
264            ),
265            "T · a.md · H1 › H2 · paragraph"
266        );
267        assert_eq!(
268            context_prefix("title", "path", &[], "paragraph"),
269            "title · path ·  · paragraph"
270        );
271        assert_eq!(embed_input("ctx", "text"), "ctx\ntext");
272        assert_eq!(hex(&ctx_hash("ctx")), hex(&sha256("ctx")));
273        assert_eq!(
274            hex(&sha256("")),
275            "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
276        );
277    }
278
279    #[test]
280    fn doc_header_and_budget() {
281        assert_eq!(doc_header("T", "a.md", None, None), "T · a.md");
282        assert_eq!(
283            doc_header("T", "a.md", Some("note"), Some("canon")),
284            "T · a.md · type: note · layer: canon"
285        );
286        assert_eq!(doc_input("h", "body\n"), "h\nbody\n");
287        assert_eq!(token_budget(None), 496);
288        assert_eq!(token_budget(Some(64)), 48);
289        assert_eq!(token_budget(Some(16)), 1);
290        assert_eq!(token_budget(Some(0)), 1);
291        assert_eq!(DocEmbedMethod::for_input("a b c", 4), DocEmbedMethod::Whole);
292        assert_eq!(
293            DocEmbedMethod::for_input("a b c", 3),
294            DocEmbedMethod::Pooled
295        );
296        assert_eq!(DocEmbedMethod::parse("whole"), Some(DocEmbedMethod::Whole));
297        assert_eq!(DocEmbedMethod::parse("x"), None);
298    }
299
300    fn r(hash: &str, tokens: u64) -> DocEmbedBlockRef {
301        DocEmbedBlockRef {
302            content_hash: hash.to_owned(),
303            ctx: "c".to_owned(),
304            tokens,
305        }
306    }
307
308    #[test]
309    fn pooling_math() {
310        let refs = [r("a", 3), r("miss", 100), r("b", 1), r("zero", 0)];
311        let lookup = |x: &DocEmbedBlockRef| match x.content_hash.as_str() {
312            "a" => Some(vec![1.0f32, 0.0]),
313            "b" => Some(vec![0.0f32, 1.0, 9.0]), // the third component is beyond dim
314            "zero" => Some(vec![0.0f32, 0.0]),
315            _ => None,
316        };
317        // acc = (3·1 + 0, 0 + 1·1) / 5 = (0.6, 0.2); normalized.
318        let v = pool_block_vectors(2, &refs, lookup).unwrap();
319        let norm = (0.6f64 * 0.6 + 0.2 * 0.2).sqrt();
320        assert_eq!(v, vec![(0.6 / norm) as f32, (0.2 / norm) as f32]);
321        // Nothing cached → stays queued.
322        assert_eq!(pool_block_vectors(2, &refs, |_| None), None);
323        // All-zero contributions → zeros, not None.
324        assert_eq!(
325            pool_block_vectors(2, &[r("zero", 2)], |_| Some(vec![0.0, 0.0])),
326            Some(vec![0.0, 0.0])
327        );
328        // A short cached vector leaves the tail at zero.
329        assert_eq!(
330            pool_block_vectors(3, &[r("a", 1)], |_| Some(vec![2.0])),
331            Some(vec![1.0, 0.0, 0.0])
332        );
333    }
334}