Files
poimen-memory/crates/mem-llm/src/embeddings.rs
T

159 lines
4.8 KiB
Rust
Raw Normal View History

use anyhow::{anyhow, Result};
use pgvector::Vector;
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::env;
use std::time::Duration;
/// Embeddings dimensionality — must match schema and HNSW index
const EMBEDDINGS_DIM: usize = 768;
/// Maximum batch size for embeddings API (gateway limit: 32)
const BATCH_SIZE: usize = 32;
/// Embeddings client for TEI/Ollama via api.riotpiao.com gateway
/// Batches requests at ≤32 texts per call, preserves input order
#[derive(Clone)]
pub struct EmbeddingsClient {
base_url: String,
model: String,
api_key: String,
http: Client,
}
#[derive(Debug, Serialize)]
struct EmbeddingRequest {
model: String,
input: Vec<String>,
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
enum EmbeddingResponse {
Success {
#[serde(default)]
object: String,
data: Vec<EmbeddingData>,
#[serde(default)]
usage: serde_json::Value,
},
Error {
error: serde_json::Value,
},
}
#[derive(Debug, Deserialize)]
struct EmbeddingData {
embedding: Vec<f32>,
#[serde(default)]
index: usize,
}
impl EmbeddingsClient {
/// Create from environment
/// Uses api.riotpiao.com gateway (nomic-ai/nomic-embed-text-v2-moe model, 768-dim)
pub fn from_env() -> Result<Self> {
let base_url = env::var("LLM_API_BASE")
.unwrap_or_else(|_| "https://api.riotpiao.com".to_string());
let model = "nomic-ai/nomic-embed-text-v2-moe".to_string();
let api_key = env::var("LLM_API_KEY")
.unwrap_or_else(|_| String::new());
let http = Client::builder()
.timeout(Duration::from_secs(30))
.build()?;
Ok(Self {
base_url,
model,
api_key,
http,
})
}
/// Embed a single text string, returning a 768-dim vector
pub async fn embed_one(&self, text: &str) -> Result<Vector> {
let embeddings = self.embed(&[text.to_string()]).await?;
Ok(embeddings
.into_iter()
.next()
.ok_or_else(|| anyhow!("empty embedding response"))?)
}
/// Embed multiple texts, batched at ≤32 per request, preserving input order
/// Returns exactly N vectors for N input texts, each 768-dim
///
/// **Batching:** Splits input into chunks of ≤32, processes each via POST /v1/embeddings
/// **Order:** Preserves input order across batch boundaries
/// **Assertion:** Every returned vector must be exactly 768-dim, else errors loudly with model name
/// **Headers:** Sends apikey even though route currently doesn't require auth (future-proofing)
pub async fn embed(&self, texts: &[String]) -> Result<Vec<Vector>> {
if texts.is_empty() {
return Ok(Vec::new());
}
let mut all_vectors = Vec::new();
// Split into batches of ≤32
for batch in texts.chunks(BATCH_SIZE) {
let batch_vecs = self.embed_batch_internal(batch).await?;
all_vectors.extend(batch_vecs);
}
Ok(all_vectors)
}
/// Internal: embed a single batch of ≤32 texts
async fn embed_batch_internal(&self, texts: &[String]) -> Result<Vec<Vector>> {
let req = EmbeddingRequest {
model: self.model.clone(),
input: texts.to_vec(),
};
let url = format!("{}/v1/embeddings", self.base_url);
let mut builder = self.http.post(&url);
// Send apikey header even though route currently doesn't require auth
// This future-proofs for when the route's auth plugin gets enabled
if !self.api_key.is_empty() {
builder = builder.header("apikey", &self.api_key);
}
let resp = builder.json(&req).send().await?;
let _status = resp.status();
let body: EmbeddingResponse = resp.json().await?;
match body {
EmbeddingResponse::Error { error } => {
Err(anyhow!("embeddings API error: {}", error))
}
EmbeddingResponse::Success { data, .. } => {
let mut vectors: Vec<Vector> = Vec::new();
for item in data {
// Assert exactly 768 dimensions
if item.embedding.len() != EMBEDDINGS_DIM {
return Err(anyhow!(
"model {} returned {}-dim vector, expected {}",
self.model,
item.embedding.len(),
EMBEDDINGS_DIM
));
}
vectors.push(Vector::from(item.embedding));
}
Ok(vectors)
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_batch_size_constant() {
assert_eq!(BATCH_SIZE, 32);
assert_eq!(EMBEDDINGS_DIM, 768);
}
}