Skip to main content

weavatrix_memory/context/
token.rs

1use crate::{
2    domain::{MemoryFact, MemoryNode},
3    error::{MemoryError, Result},
4};
5
6pub trait TokenEstimator {
7    fn estimate(&self, value: &str) -> usize;
8    fn name(&self) -> &'static str;
9}
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
12pub struct BytesTokenEstimator {
13    bytes_per_token: usize,
14}
15
16impl BytesTokenEstimator {
17    /// Creates a deterministic byte-based fallback estimator.
18    ///
19    /// # Errors
20    ///
21    /// Rejects zero bytes per token.
22    pub fn new(bytes_per_token: usize) -> Result<Self> {
23        if bytes_per_token == 0 {
24            return Err(MemoryError::InvalidValue {
25                field: "bytes_per_token",
26                reason: "must be greater than zero",
27            });
28        }
29        Ok(Self { bytes_per_token })
30    }
31}
32
33impl Default for BytesTokenEstimator {
34    fn default() -> Self {
35        Self { bytes_per_token: 4 }
36    }
37}
38
39impl TokenEstimator for BytesTokenEstimator {
40    fn estimate(&self, value: &str) -> usize {
41        value.len().div_ceil(self.bytes_per_token).max(1)
42    }
43
44    fn name(&self) -> &'static str {
45        "utf8_bytes"
46    }
47}
48
49pub(super) fn node_tokens(estimator: &impl TokenEstimator, node: &MemoryNode) -> usize {
50    estimator.estimate(node.id.as_str())
51        + estimator.estimate(&node.kind)
52        + estimator.estimate(&node.label)
53        + node
54            .repository
55            .iter()
56            .chain(node.branch.iter())
57            .map(|value| estimator.estimate(value))
58            .sum::<usize>()
59        + node
60            .attributes
61            .iter()
62            .map(|(key, value)| estimator.estimate(key) + estimator.estimate(value))
63            .sum::<usize>()
64        + 6
65}
66
67pub(super) fn fact_tokens(estimator: &impl TokenEstimator, fact: &MemoryFact) -> usize {
68    estimator.estimate(fact.id.as_str())
69        + estimator.estimate(fact.source.as_str())
70        + estimator.estimate(&fact.relation)
71        + estimator.estimate(fact.target.as_str())
72        + fact
73            .evidence
74            .iter()
75            .map(|item| {
76                estimator.estimate(&item.kind)
77                    + estimator.estimate(&item.source)
78                    + item
79                        .locator
80                        .iter()
81                        .chain(item.digest.iter())
82                        .map(|value| estimator.estimate(value))
83                        .sum::<usize>()
84            })
85            .sum::<usize>()
86        + 16
87}