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 /// Uses api.riotpiao.com gateway (nomic-ai/nomic-embed-text-v2-moe model, 768-dim) pub fn from_env() -> Result { 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 { 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); } }