embed.rsannotatedembed.rssource106 lines · 3.6 KB · raw
1//! Turning text into vectors, so Whiskers can search what it remembers by meaning.
2//!
3//! This is its own port rather than a method of the judge: it is not Jev (today a local model server
4//! on the operator's machine), nothing she says should leave that machine for it, and a hosted backend
5//! has no such machine, so whether it is available differs from whether Jev is.
6
7use std::fmt;
8
9use crate::error::Diagnostic;
10
11/// The most texts in one batch.
12pub const EMBED_TEXTS_MAX: usize = 64;
13
14/// The longest text, in characters.
15pub const EMBED_CHARS_MAX: usize = 2000;
16
17/// Texts to embed: between one and [`EMBED_TEXTS_MAX`], none longer than [`EMBED_CHARS_MAX`].
18#[derive(Clone, Debug, PartialEq, Eq)]
19pub struct EmbedBatch(Vec<String>);
20
21#[derive(Clone, Copy, Debug, PartialEq, Eq)]
22pub enum BatchError {
23    Count(usize),
24    TooLong,
25}
26
27impl fmt::Display for BatchError {
28    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
29        match self {
30            BatchError::Count(n) => write!(f, "1 to {EMBED_TEXTS_MAX} texts, not {n}"),
31            BatchError::TooLong => write!(f, "a text is longer than {EMBED_CHARS_MAX} characters"),
32        }
33    }
34}
35
36impl std::error::Error for BatchError {}
37
38impl EmbedBatch {
39    pub fn new(texts: Vec<String>) -> Result<Self, BatchError> {
40        if texts.is_empty() || texts.len() > EMBED_TEXTS_MAX {
41            ::log::warn!("embed refused: {} texts (1 to {EMBED_TEXTS_MAX} allowed)", texts.len());
42            return Err(BatchError::Count(texts.len()));
43        }
44        if texts.iter().any(|t| t.chars().count() > EMBED_CHARS_MAX) {
45            ::log::warn!("embed refused: a text is longer than {EMBED_CHARS_MAX} characters");
46            return Err(BatchError::TooLong);
47        }
48        Ok(Self(texts))
49    }
50
51    pub fn texts(&self) -> &[String] {
52        &self.0
53    }
54
55    pub fn len(&self) -> usize {
56        self.0.len()
57    }
58
59    /// A batch is never empty; this exists so `len` has its usual partner.
60    pub fn is_empty(&self) -> bool {
61        false
62    }
63}
64
65/// Why no vectors came back.
66#[derive(Clone, Debug, PartialEq, Eq)]
67pub enum EmbedError {
68    Unreachable(Diagnostic),
69    Refused { status: u16, said: Diagnostic },
70    Unreadable(Diagnostic),
71    /// The server returned a different number of vectors than texts.
72    WrongCount { asked: usize, got: usize },
73}
74
75impl fmt::Display for EmbedError {
76    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
77        match self {
78            EmbedError::Unreachable(d) | EmbedError::Unreadable(d) | EmbedError::Refused { said: d, .. } => d.fmt(f),
79            EmbedError::WrongCount { asked, got } => write!(f, "embedding server: asked for {asked} vectors, got {got}"),
80        }
81    }
82}
83
84impl std::error::Error for EmbedError {}
85
86/// The embedding model.
87///
88/// Contract: one vector per text, in the texts' order, or an `Err`; never a partial answer.
89#[expect(async_fn_in_trait, reason = "a Worker's futures hold JavaScript values and cannot be Send, so no Send bound may be required here")]
90pub trait Embedder {
91    async fn embed(&self, batch: &EmbedBatch) -> Result<Vec<Vec<f32>>, EmbedError>;
92}
93
94#[cfg(test)]
95mod tests {
96    use super::*;
97
98    #[test]
99    fn a_batch_is_one_to_sixty_four_texts_of_a_sane_length() {
100        assert!(EmbedBatch::new(vec!["a".into()]).is_ok());
101        assert_eq!(EmbedBatch::new(vec![]), Err(BatchError::Count(0)));
102        assert_eq!(EmbedBatch::new(vec!["x".into(); 65]), Err(BatchError::Count(65)));
103        assert_eq!(EmbedBatch::new(vec!["x".repeat(EMBED_CHARS_MAX + 1)]), Err(BatchError::TooLong));
104        assert_eq!(BatchError::Count(0).to_string(), "1 to 64 texts, not 0");
105    }
106}