use sha2::{Digest, Sha256};
use crate::fts::split_ws;
pub const MIN_EMBED_TOKENS: usize = 24;
pub const EMBED_BATCH: usize = 32;
pub const DEFAULT_DOC_TOKEN_BUDGET: u32 = 512;
pub const DOC_HEADER_MARGIN_TOKENS: u32 = 16;
#[must_use]
pub fn sha256(s: &str) -> [u8; 32] {
Sha256::digest(s.as_bytes()).into()
}
#[must_use]
pub fn hex(bytes: &[u8]) -> String {
bytes.iter().map(|b| format!("{b:02x}")).collect()
}
#[must_use]
pub fn word_count(text: &str) -> usize {
split_ws(text).count()
}
#[must_use]
pub fn should_embed(text: &str) -> bool {
word_count(text) >= MIN_EMBED_TOKENS
}
#[must_use]
pub fn estimate_tokens(text: &str) -> u64 {
(word_count(text) as f64 * 1.3).ceil() as u64
}
#[must_use]
pub fn context_prefix(
doc_title: &str,
path: &str,
heading_chain: &[String],
block_type: &str,
) -> String {
format!(
"{doc_title} \u{00B7} {path} \u{00B7} {} \u{00B7} {block_type}",
heading_chain.join(" \u{203A} ")
)
}
#[must_use]
pub fn embed_input(ctx: &str, block_text: &str) -> String {
format!("{ctx}\n{block_text}")
}
#[must_use]
pub fn ctx_hash(ctx: &str) -> [u8; 32] {
sha256(ctx)
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct EmbedTask {
pub block_id: String,
pub content_hash: String,
pub ctx: String,
pub text: String,
}
impl EmbedTask {
#[must_use]
pub fn input(&self) -> String {
embed_input(&self.ctx, &self.text)
}
}
#[must_use]
pub fn doc_header(title: &str, path: &str, doc_type: Option<&str>, layer: Option<&str>) -> String {
let mut bits = vec![title.to_owned(), path.to_owned()];
if let Some(t) = doc_type {
bits.push(format!("type: {t}"));
}
if let Some(l) = layer {
bits.push(format!("layer: {l}"));
}
bits.join(" \u{00B7} ")
}
#[must_use]
pub fn doc_input(header: &str, body: &str) -> String {
format!("{header}\n{body}")
}
#[must_use]
pub fn token_budget(max_input_tokens: Option<u32>) -> u64 {
let budget = max_input_tokens.unwrap_or(DEFAULT_DOC_TOKEN_BUDGET);
u64::from(budget.saturating_sub(DOC_HEADER_MARGIN_TOKENS).max(1))
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DocEmbedMethod {
Whole,
Pooled,
}
impl DocEmbedMethod {
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
DocEmbedMethod::Whole => "whole",
DocEmbedMethod::Pooled => "pooled",
}
}
#[must_use]
pub fn parse(s: &str) -> Option<Self> {
match s {
"whole" => Some(DocEmbedMethod::Whole),
"pooled" => Some(DocEmbedMethod::Pooled),
_ => None,
}
}
#[must_use]
pub fn for_input(input: &str, budget: u64) -> Self {
if estimate_tokens(input) <= budget {
DocEmbedMethod::Whole
} else {
DocEmbedMethod::Pooled
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct DocEmbedBlockRef {
pub content_hash: String,
pub ctx: String,
pub tokens: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct DocEmbedTask {
pub doc_id: String,
pub header: String,
pub input: String,
pub blocks: Vec<DocEmbedBlockRef>,
}
impl DocEmbedTask {
#[must_use]
pub fn input_hash(&self) -> [u8; 32] {
sha256(&self.input)
}
}
pub fn pool_block_vectors<F>(
dim: usize,
refs: &[DocEmbedBlockRef],
mut cached: F,
) -> Option<Vec<f32>>
where
F: FnMut(&DocEmbedBlockRef) -> Option<Vec<f32>>,
{
let mut acc = vec![0.0f64; dim];
let mut weight_sum = 0.0f64;
for r in refs {
let Some(v) = cached(r) else {
continue;
};
let w = if r.tokens > 0 { r.tokens as f64 } else { 1.0 };
let n = dim.min(v.len());
for i in 0..n {
acc[i] += w * f64::from(v[i]);
}
weight_sum += w;
}
if weight_sum == 0.0 {
return None;
}
let mut norm = 0.0f64;
for a in &mut acc {
*a /= weight_sum;
norm += *a * *a;
}
let norm = norm.sqrt();
if norm == 0.0 {
return Some(vec![0.0f32; dim]);
}
Some(acc.iter().map(|a| (a / norm) as f32).collect())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn should_embed_threshold_and_estimate() {
let words = |n: usize| {
(0..n)
.map(|i| format!("w{i}"))
.collect::<Vec<_>>()
.join(" ")
};
assert!(!should_embed(&words(23)));
assert!(should_embed(&words(24)));
assert!(should_embed(&format!(" {} \n", words(24))));
assert!(!should_embed(""));
assert_eq!(estimate_tokens(""), 0);
assert_eq!(estimate_tokens("one"), 2);
assert_eq!(estimate_tokens("one two three"), 4); assert_eq!(estimate_tokens(&words(10)), 13);
assert_eq!(estimate_tokens(&words(24)), 32); }
#[test]
fn context_prefix_shape() {
assert_eq!(
context_prefix(
"T",
"a.md",
&["H1".to_owned(), "H2".to_owned()],
"paragraph"
),
"T · a.md · H1 › H2 · paragraph"
);
assert_eq!(
context_prefix("title", "path", &[], "paragraph"),
"title · path · · paragraph"
);
assert_eq!(embed_input("ctx", "text"), "ctx\ntext");
assert_eq!(hex(&ctx_hash("ctx")), hex(&sha256("ctx")));
assert_eq!(
hex(&sha256("")),
"e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
);
}
#[test]
fn doc_header_and_budget() {
assert_eq!(doc_header("T", "a.md", None, None), "T · a.md");
assert_eq!(
doc_header("T", "a.md", Some("note"), Some("canon")),
"T · a.md · type: note · layer: canon"
);
assert_eq!(doc_input("h", "body\n"), "h\nbody\n");
assert_eq!(token_budget(None), 496);
assert_eq!(token_budget(Some(64)), 48);
assert_eq!(token_budget(Some(16)), 1);
assert_eq!(token_budget(Some(0)), 1);
assert_eq!(DocEmbedMethod::for_input("a b c", 4), DocEmbedMethod::Whole);
assert_eq!(
DocEmbedMethod::for_input("a b c", 3),
DocEmbedMethod::Pooled
);
assert_eq!(DocEmbedMethod::parse("whole"), Some(DocEmbedMethod::Whole));
assert_eq!(DocEmbedMethod::parse("x"), None);
}
fn r(hash: &str, tokens: u64) -> DocEmbedBlockRef {
DocEmbedBlockRef {
content_hash: hash.to_owned(),
ctx: "c".to_owned(),
tokens,
}
}
#[test]
fn pooling_math() {
let refs = [r("a", 3), r("miss", 100), r("b", 1), r("zero", 0)];
let lookup = |x: &DocEmbedBlockRef| match x.content_hash.as_str() {
"a" => Some(vec![1.0f32, 0.0]),
"b" => Some(vec![0.0f32, 1.0, 9.0]), "zero" => Some(vec![0.0f32, 0.0]),
_ => None,
};
let v = pool_block_vectors(2, &refs, lookup).unwrap();
let norm = (0.6f64 * 0.6 + 0.2 * 0.2).sqrt();
assert_eq!(v, vec![(0.6 / norm) as f32, (0.2 / norm) as f32]);
assert_eq!(pool_block_vectors(2, &refs, |_| None), None);
assert_eq!(
pool_block_vectors(2, &[r("zero", 2)], |_| Some(vec![0.0, 0.0])),
Some(vec![0.0, 0.0])
);
assert_eq!(
pool_block_vectors(3, &[r("a", 1)], |_| Some(vec![2.0])),
Some(vec![1.0, 0.0, 0.0])
);
}
}