whiskers.git / crates / whiskersd / src / embed.rs
embed.rsannotatedembed.rssource138 lines · 5.2 KB · raw

Turning text into vectors, by a local llama.cpp server running an embedding model on this machine (OpenAI-compatible /v1/embeddings). Nothing she says leaves the machine for this.

4use std::time::Duration;
6use log::{debug, trace, warn};
7use whiskers_ports::{Diagnostic, EmbedBatch, EmbedError, Embedder};
8
9const DEFAULT_URL: &str = "http://127.0.0.1:8081";
10
11pub struct LlamaEmbedder {
12    url: String,
13    agent: ureq::Agent,
14}
15
16impl LlamaEmbedder {

WHISKERS_EMBED_URL overrides where the llama.cpp server listens.

18    pub fn from_env() -> Self {
19        let url = std::env::var("WHISKERS_EMBED_URL").ok().filter(|v| !v.trim().is_empty()).unwrap_or_else(|| DEFAULT_URL.to_owned());
20        debug!("embedding server at {url}");
21        Self::new(url)
22    }
24    pub fn new(url: String) -> Self {
25        let agent = ureq::Agent::config_builder()
26            .timeout_global(Some(Duration::from_secs(10)))
27            .http_status_as_error(false)
28            .build()
29            .into();
30        Self { url: url.trim_end_matches('/').to_owned(), agent }
31    }
32
33    fn ask(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, EmbedError> {
34        trace!("embed: {} text(s)", texts.len());
35        let body = serde_json::json!({ "input": texts, "model": "embedding" }).to_string();
36        let mut resp = self
37            .agent
38            .post(format!("{}/v1/embeddings", self.url))
39            .header("content-type", "application/json")
40            .send(body)
41            .map_err(|e| {
42                warn!("embedding server unreachable: {e}");
43                EmbedError::Unreachable(Diagnostic::new(format!("embedding server: {e}")))
44            })?;
45        let status = resp.status();
46        let text = resp.body_mut().read_to_string().map_err(|e| {
47            warn!("embedding server reply unreadable (status {status}): {e}");
48            EmbedError::Unreadable(Diagnostic::new(format!("embedding server: {e}")))
49        })?;
50        debug!("embedding server answered {status}, {} bytes", text.len());
51        if !status.is_success() {
52            warn!("embedding server refused with {status}");
53            return Err(EmbedError::Refused {
54                status: status.as_u16(),
55                said: Diagnostic::new(format!("embedding server: {status}: {}", text.chars().take(160).collect::<String>())),
56            });
57        }
58        vectors_of(&text, texts.len())
59    }
60}
61
62impl Embedder for LlamaEmbedder {
63    async fn embed(&self, batch: &EmbedBatch) -> Result<Vec<Vec<f32>>, EmbedError> {
64        self.ask(batch.texts())
65    }
66}

The vectors in an OpenAI-style embeddings reply, in input order.

69fn vectors_of(body: &str, expected: usize) -> Result<Vec<Vec<f32>>, EmbedError> {
70    #[derive(serde::Deserialize)]
71    struct Item {
72        index: usize,
73        embedding: Vec<f32>,
74    }
75    #[derive(serde::Deserialize)]
76    struct Reply {
77        data: Vec<Item>,
78    }
79    let mut r: Reply = serde_json::from_str(body).map_err(|e| {
80        warn!("embedding reply ({} bytes) does not parse: {e}", body.len());
81        EmbedError::Unreadable(Diagnostic::new(format!("embedding server: {e}")))
82    })?;
83    if r.data.len() != expected {
84        warn!("embedding reply has {} vectors, expected {expected}", r.data.len());
85        return Err(EmbedError::WrongCount { asked: expected, got: r.data.len() });
86    }
87    r.data.sort_by_key(|i| i.index);
88    Ok(r.data.into_iter().map(|i| i.embedding).collect())
89}
91#[cfg(test)]
92mod tests {
93    use std::io::{Read, Write};
94    use std::net::TcpListener;
95    use std::thread;
96
97    use super::*;
98    use whiskers_ports::run_ready;
99
100    fn upstream(body: &str) -> String {
101        let l = TcpListener::bind("127.0.0.1:0").unwrap();
102        let url = format!("http://{}", l.local_addr().unwrap());
103        let body = body.to_owned();
104        thread::spawn(move || {
105            let (mut s, _) = l.accept().unwrap();
106            let mut buf = [0u8; 65536];
107            let _ = s.read(&mut buf);
108            let r = format!("HTTP/1.1 200 OK\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}", body.len());
109            s.write_all(r.as_bytes()).unwrap();
110        });
111        url
112    }
113
114    fn batch(texts: &[&str]) -> EmbedBatch {
115        EmbedBatch::new(texts.iter().map(|t| (*t).to_owned()).collect()).unwrap()
116    }
117
118    #[test]
119    fn vectors_come_back_in_input_order_whatever_order_the_server_used() {
120        let url = upstream(r#"{"data":[{"index":1,"embedding":[3,4]},{"index":0,"embedding":[1,2]}]}"#);
121        let got = run_ready(LlamaEmbedder::new(url).embed(&batch(&["a", "b"]))).unwrap();
122        assert_eq!(got, vec![vec![1.0, 2.0], vec![3.0, 4.0]]);
123    }
124
125    #[test]
126    fn a_wrong_count_or_nonsense_is_an_error() {
127        let url = upstream(r#"{"data":[{"index":0,"embedding":[1]}]}"#);
128        assert_eq!(run_ready(LlamaEmbedder::new(url).embed(&batch(&["a", "b"]))), Err(EmbedError::WrongCount { asked: 2, got: 1 }));
129        let url = upstream("nonsense");
130        assert!(matches!(run_ready(LlamaEmbedder::new(url).embed(&batch(&["a"]))), Err(EmbedError::Unreadable(_))));
131    }
132
133    #[test]
134    fn a_server_that_is_not_there_is_unreachable() {
135        let s = LlamaEmbedder::new("http://127.0.0.1:1".into());
136        assert!(matches!(run_ready(s.embed(&batch(&["x"]))), Err(EmbedError::Unreachable(_))));
137    }
138}