1//! Turning text into vectors, by a local llama.cpp server running an embedding model on this 2//! machine (OpenAI-compatible `/v1/embeddings`). Nothing she says leaves the machine for this. 3 4use std::time::Duration; 5 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 { 17 /// `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 } 23 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} 67 68/// 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} 90 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}