weavatrix_memory/context/
token.rs1use 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 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}