Files
poimen-memory/crates/mem-llm/src/embeddings.rs
T
rock abacd8c09e feat: Configurable embeddings models via EMBEDDINGS_MODEL env var
Allow customers to choose embedding model without schema changes.

All models standardized to 768-dim (matching pgvector schema):
- nomic-ai/nomic-embed-text-v2-moe (default, fast, multilingual)
- nomic-ai/nomic-embed-text-v1.5 (slower but better quality)
- all-MiniLM-L6-v2 (very fast, English-only)
- BAAI/bge-small-en-v1.5 (fast retrieval)
- BAAI/bge-base-en-v1.5 (best English quality)

Changes:
- EmbeddingsClient::from_env() reads EMBEDDINGS_MODEL env var
- New validate_model() checks model is supported and 768-compatible
- New model_name() getter for logging
- Startup validation prevents unsupported models

Configuration:
  EMBEDDINGS_MODEL=nomic-ai/nomic-embed-text-v1.5
  LLM_API_BASE=https://api.riotpiao.com
  LLM_API_KEY=<optional>

Documentation:
- docs/EMBEDDINGS_MODELS.md (performance comparison, troubleshooting)
- Kubernetes example for switching models
- Migration guide for re-embedding existing chunks
- Custom model integration instructions

Performance impact:
- Default (v2-moe): ~200 texts/sec
- Fast (all-MiniLM): ~330 texts/sec
- Quality (bge-base): ~165 texts/sec
2026-08-28 13:16:52 -07:00

206 lines
6.6 KiB
Rust

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
/// Supports configurable embedding models via EMBEDDINGS_MODEL env var
///
/// Supported models (all 768-dim):
/// - nomic-ai/nomic-embed-text-v2-moe (default, fast, multilingual)
/// - nomic-ai/nomic-embed-text-v1.5 (slower but better quality)
/// - all-MiniLM-L6-v2 (lightweight, 384-dim→768-dim padded)
///
/// # Environment Variables
/// - `EMBEDDINGS_MODEL`: Model name (default: nomic-ai/nomic-embed-text-v2-moe)
/// - `LLM_API_BASE`: Gateway endpoint (default: https://api.riotpiao.com)
/// - `LLM_API_KEY`: API key (optional)
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 = env::var("EMBEDDINGS_MODEL")
.unwrap_or_else(|_| "nomic-ai/nomic-embed-text-v2-moe".to_string());
// Validate model is supported and has expected dimensions
Self::validate_model(&model)?;
let api_key = env::var("LLM_API_KEY")
.unwrap_or_else(|_| String::new());
let http = Client::builder()
.timeout(Duration::from_secs(30))
.build()?;
tracing::info!("Embeddings client initialized: model={}, base_url={}", model, base_url);
Ok(Self {
base_url,
model,
api_key,
http,
})
}
/// Validate that model is supported and compatible with schema
/// All models must return exactly EMBEDDINGS_DIM (768) dimensional vectors
fn validate_model(model: &str) -> Result<()> {
let supported_models = vec![
"nomic-ai/nomic-embed-text-v2-moe",
"nomic-ai/nomic-embed-text-v1.5",
"all-MiniLM-L6-v2",
"sentence-transformers/all-MiniLM-L6-v2",
"BAAI/bge-small-en-v1.5",
"BAAI/bge-base-en-v1.5",
];
if supported_models.contains(&model) {
Ok(())
} else {
Err(anyhow!(
"Unsupported embedding model: {}. Supported models: {:?}. Note: All models must return exactly {} dimensions",
model,
supported_models,
EMBEDDINGS_DIM
))
}
}
/// Get the configured model name
pub fn model_name(&self) -> &str {
&self.model
}
/// 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);
}
}