1use sha2::{Digest, Sha256};
6
7use crate::fts::split_ws;
8
9pub const MIN_EMBED_TOKENS: usize = 24;
11pub const EMBED_BATCH: usize = 32;
13pub const DEFAULT_DOC_TOKEN_BUDGET: u32 = 512;
15pub const DOC_HEADER_MARGIN_TOKENS: u32 = 16;
17
18#[must_use]
20pub fn sha256(s: &str) -> [u8; 32] {
21 Sha256::digest(s.as_bytes()).into()
22}
23
24#[must_use]
26pub fn hex(bytes: &[u8]) -> String {
27 bytes.iter().map(|b| format!("{b:02x}")).collect()
28}
29
30#[must_use]
32pub fn word_count(text: &str) -> usize {
33 split_ws(text).count()
34}
35
36#[must_use]
39pub fn should_embed(text: &str) -> bool {
40 word_count(text) >= MIN_EMBED_TOKENS
41}
42
43#[must_use]
45pub fn estimate_tokens(text: &str) -> u64 {
46 (word_count(text) as f64 * 1.3).ceil() as u64
47}
48
49#[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#[must_use]
65pub fn embed_input(ctx: &str, block_text: &str) -> String {
66 format!("{ctx}\n{block_text}")
67}
68
69#[must_use]
71pub fn ctx_hash(ctx: &str) -> [u8; 32] {
72 sha256(ctx)
73}
74
75#[derive(Clone, Debug, PartialEq, Eq)]
77pub struct EmbedTask {
78 pub block_id: String,
79 pub content_hash: String,
81 pub ctx: String,
82 pub text: String,
83}
84
85impl EmbedTask {
86 #[must_use]
88 pub fn input(&self) -> String {
89 embed_input(&self.ctx, &self.text)
90 }
91}
92
93#[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#[must_use]
110pub fn doc_input(header: &str, body: &str) -> String {
111 format!("{header}\n{body}")
112}
113
114#[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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
123pub enum DocEmbedMethod {
124 Whole,
126 Pooled,
128}
129
130impl DocEmbedMethod {
131 #[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 #[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 #[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#[derive(Clone, Debug, PartialEq, Eq)]
164pub struct DocEmbedBlockRef {
165 pub content_hash: String,
167 pub ctx: String,
168 pub tokens: u64,
170}
171
172#[derive(Clone, Debug, PartialEq, Eq)]
174pub struct DocEmbedTask {
175 pub doc_id: String,
176 pub header: String,
178 pub input: String,
180 pub blocks: Vec<DocEmbedBlockRef>,
182}
183
184impl DocEmbedTask {
185 #[must_use]
187 pub fn input_hash(&self) -> [u8; 32] {
188 sha256(&self.input)
189 }
190}
191
192pub 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); assert_eq!(estimate_tokens(&words(10)), 13);
253 assert_eq!(estimate_tokens(&words(24)), 32); }
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]), "zero" => Some(vec![0.0f32, 0.0]),
315 _ => None,
316 };
317 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 assert_eq!(pool_block_vectors(2, &refs, |_| None), None);
323 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 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}