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, } #[derive(Debug, Deserialize)] #[serde(untagged)] enum EmbeddingResponse { Success { #[serde(default)] object: String, data: Vec, #[serde(default)] usage: serde_json::Value, }, Error { error: serde_json::Value, }, } #[derive(Debug, Deserialize)] struct EmbeddingData { embedding: Vec, #[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 { 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 { 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> { 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> { 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 = 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); } }