Merge commit 'd803bfe2b1fe7f5e219e50ac20d6801a0a58ac75' as 'vendor/ruvector'

This commit is contained in:
ruv
2026-02-28 14:39:40 -05:00
7854 changed files with 3522914 additions and 0 deletions
@@ -0,0 +1,152 @@
//! Thread-safe model caching with lazy loading
use dashmap::DashMap;
use fastembed::{EmbeddingModel as FastEmbedModel, InitOptions, TextEmbedding};
use parking_lot::RwLock;
use super::models::EmbeddingModel;
/// Global model cache for lazy loading and reuse
pub struct ModelCache {
/// Cached embedding models (using RwLock for interior mutability)
models: DashMap<EmbeddingModel, RwLock<TextEmbedding>>,
/// Default model setting
default_model: RwLock<EmbeddingModel>,
}
impl ModelCache {
/// Create a new model cache
pub fn new() -> Self {
Self {
models: DashMap::new(),
default_model: RwLock::new(EmbeddingModel::default()),
}
}
/// Get or load a model and generate embeddings
pub fn embed(&self, model: EmbeddingModel, texts: Vec<&str>) -> Result<Vec<Vec<f32>>, String> {
// Check if already cached
if let Some(cached) = self.models.get(&model) {
let mut embedding = cached.write();
return embedding
.embed(texts, None)
.map_err(|e| format!("Embedding failed: {}", e));
}
// Load the model
let embedding = self.load_model(model)?;
// Generate embeddings first
let mut embedding_model = embedding;
let result = embedding_model
.embed(texts, None)
.map_err(|e| format!("Embedding failed: {}", e));
// Cache the model
self.models.insert(model, RwLock::new(embedding_model));
result
}
/// Load a model from fastembed
fn load_model(&self, model: EmbeddingModel) -> Result<TextEmbedding, String> {
let fastembed_model = match model {
EmbeddingModel::AllMiniLmL6V2 => FastEmbedModel::AllMiniLML6V2,
EmbeddingModel::BgeSmallEnV15 => FastEmbedModel::BGESmallENV15,
EmbeddingModel::BgeBaseEnV15 => FastEmbedModel::BGEBaseENV15,
EmbeddingModel::BgeLargeEnV15 => FastEmbedModel::BGELargeENV15,
EmbeddingModel::AllMpnetBaseV2 => FastEmbedModel::AllMiniLML6V2, // Fallback
EmbeddingModel::NomicEmbedTextV15 => FastEmbedModel::NomicEmbedTextV15,
};
let options = InitOptions::new(fastembed_model).with_show_download_progress(false);
TextEmbedding::try_new(options)
.map_err(|e| format!("Failed to load model '{}': {}", model.name(), e))
}
/// Pre-load a model into the cache
pub fn preload(&self, model: EmbeddingModel) -> Result<(), String> {
if self.models.contains_key(&model) {
return Ok(());
}
let embedding = self.load_model(model)?;
self.models.insert(model, RwLock::new(embedding));
Ok(())
}
/// Check if a model is loaded
pub fn is_loaded(&self, model: EmbeddingModel) -> bool {
self.models.contains_key(&model)
}
/// Get list of loaded models
pub fn loaded_models(&self) -> Vec<EmbeddingModel> {
self.models.iter().map(|r| *r.key()).collect()
}
/// Unload a model from cache
pub fn unload(&self, model: EmbeddingModel) -> bool {
self.models.remove(&model).is_some()
}
/// Clear all cached models
pub fn clear(&self) {
self.models.clear();
}
/// Get the default model
pub fn default_model(&self) -> EmbeddingModel {
*self.default_model.read()
}
/// Set the default model
pub fn set_default_model(&self, model: EmbeddingModel) {
*self.default_model.write() = model;
}
/// Get memory usage estimate in bytes
pub fn estimated_memory_usage(&self) -> usize {
self.models
.iter()
.map(|r| r.key().memory_mb() * 1024 * 1024)
.sum()
}
}
impl Default for ModelCache {
fn default() -> Self {
Self::new()
}
}
// Global singleton cache
lazy_static::lazy_static! {
pub static ref GLOBAL_CACHE: ModelCache = ModelCache::new();
}
/// Get the global model cache
pub fn global_cache() -> &'static ModelCache {
&GLOBAL_CACHE
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_cache_creation() {
let cache = ModelCache::new();
assert!(!cache.is_loaded(EmbeddingModel::AllMiniLmL6V2));
assert!(cache.loaded_models().is_empty());
}
#[test]
fn test_default_model() {
let cache = ModelCache::new();
assert_eq!(cache.default_model(), EmbeddingModel::AllMiniLmL6V2);
cache.set_default_model(EmbeddingModel::BgeSmallEnV15);
assert_eq!(cache.default_model(), EmbeddingModel::BgeSmallEnV15);
}
}
@@ -0,0 +1,368 @@
//! SQL function implementations for embedding generation
use pgrx::prelude::*;
use super::cache::global_cache;
use super::models::{EmbeddingModel, ModelInfo};
use super::{MAX_BATCH_SIZE, MAX_TEXT_LENGTH};
// ============================================================================
// Core Embedding Functions
// ============================================================================
/// Generate an embedding vector from text
///
/// # Arguments
/// * `text` - The text to embed
/// * `model_name` - Optional model name (defaults to 'all-MiniLM-L6-v2')
///
/// # Returns
/// A vector of f32 values representing the text embedding
///
/// # Example
/// ```sql
/// SELECT ruvector_embed('Hello world');
/// SELECT ruvector_embed('Hello world', 'bge-small');
/// ```
#[pg_extern(immutable, parallel_safe)]
pub fn ruvector_embed(text: &str, model_name: default!(&str, "'all-MiniLM-L6-v2'")) -> Vec<f32> {
// Validate text length
if text.len() > MAX_TEXT_LENGTH {
pgrx::error!(
"Text length {} exceeds maximum {} characters",
text.len(),
MAX_TEXT_LENGTH
);
}
// Parse model name
let model = EmbeddingModel::from_name(model_name).unwrap_or_else(|| {
pgrx::warning!("Unknown model '{}', using default", model_name);
EmbeddingModel::default()
});
// Generate embedding using cached model
let documents = vec![text];
match global_cache().embed(model, documents) {
Ok(embeddings) => {
if let Some(embedding) = embeddings.into_iter().next() {
embedding
} else {
pgrx::error!("No embedding generated");
}
}
Err(e) => {
pgrx::error!("Embedding generation failed: {}", e);
}
}
}
/// Generate embeddings for multiple texts in batch
///
/// # Arguments
/// * `texts` - Array of texts to embed
/// * `model_name` - Optional model name (defaults to 'all-MiniLM-L6-v2')
///
/// # Returns
/// A 2D array of embeddings (one per input text)
///
/// # Example
/// ```sql
/// SELECT ruvector_embed_batch(ARRAY['Hello', 'World', 'Test']);
/// ```
#[pg_extern(immutable, parallel_safe)]
pub fn ruvector_embed_batch(
texts: Vec<String>,
model_name: default!(&str, "'all-MiniLM-L6-v2'"),
) -> Vec<Vec<f32>> {
// Validate batch size
if texts.len() > MAX_BATCH_SIZE {
pgrx::error!(
"Batch size {} exceeds maximum {}",
texts.len(),
MAX_BATCH_SIZE
);
}
// Validate text lengths
for (i, text) in texts.iter().enumerate() {
if text.len() > MAX_TEXT_LENGTH {
pgrx::error!(
"Text at index {} exceeds maximum {} characters",
i,
MAX_TEXT_LENGTH
);
}
}
// Parse model name
let model = EmbeddingModel::from_name(model_name).unwrap_or_else(|| {
pgrx::warning!("Unknown model '{}', using default", model_name);
EmbeddingModel::default()
});
// Generate embeddings using cached model
let documents: Vec<&str> = texts.iter().map(|s| s.as_str()).collect();
match global_cache().embed(model, documents) {
Ok(embeddings) => embeddings,
Err(e) => {
pgrx::error!("Batch embedding generation failed: {}", e);
}
}
}
// ============================================================================
// Model Management Functions
// ============================================================================
/// List all available embedding models
///
/// # Returns
/// Table with model name, dimensions, and description
///
/// # Example
/// ```sql
/// SELECT * FROM ruvector_embedding_models();
/// ```
#[pg_extern]
pub fn ruvector_embedding_models() -> TableIterator<
'static,
(
name!(name, String),
name!(dimensions, i32),
name!(description, String),
name!(speed, i32),
name!(quality, i32),
name!(memory_mb, i32),
name!(loaded, bool),
),
> {
let cache = global_cache();
let rows: Vec<_> = EmbeddingModel::all()
.iter()
.map(|model| {
(
model.name().to_string(),
model.dimensions() as i32,
model.description().to_string(),
model.speed_rating() as i32,
model.quality_rating() as i32,
model.memory_mb() as i32,
cache.is_loaded(*model),
)
})
.collect();
TableIterator::new(rows)
}
/// Pre-load an embedding model into cache
///
/// # Arguments
/// * `model_name` - Name of the model to load
///
/// # Returns
/// true if loaded successfully
///
/// # Example
/// ```sql
/// SELECT ruvector_load_model('bge-small');
/// ```
#[pg_extern]
pub fn ruvector_load_model(model_name: &str) -> bool {
let model = match EmbeddingModel::from_name(model_name) {
Some(m) => m,
None => {
pgrx::warning!("Unknown model: {}", model_name);
return false;
}
};
match global_cache().preload(model) {
Ok(_) => {
pgrx::info!("Model '{}' loaded successfully", model.name());
true
}
Err(e) => {
pgrx::warning!("Failed to load model '{}': {}", model_name, e);
false
}
}
}
/// Unload an embedding model from cache
///
/// # Arguments
/// * `model_name` - Name of the model to unload
///
/// # Returns
/// true if the model was unloaded
///
/// # Example
/// ```sql
/// SELECT ruvector_unload_model('bge-small');
/// ```
#[pg_extern]
pub fn ruvector_unload_model(model_name: &str) -> bool {
let model = match EmbeddingModel::from_name(model_name) {
Some(m) => m,
None => {
pgrx::warning!("Unknown model: {}", model_name);
return false;
}
};
global_cache().unload(model)
}
/// Get information about a specific model
///
/// # Arguments
/// * `model_name` - Name of the model
///
/// # Returns
/// JSON object with model information
///
/// # Example
/// ```sql
/// SELECT ruvector_model_info('all-MiniLM-L6-v2');
/// ```
#[pg_extern]
pub fn ruvector_model_info(model_name: &str) -> pgrx::JsonB {
let model = match EmbeddingModel::from_name(model_name) {
Some(m) => m,
None => {
return pgrx::JsonB(serde_json::json!({
"error": format!("Unknown model: {}", model_name),
"available_models": EmbeddingModel::all().iter().map(|m| m.name()).collect::<Vec<_>>()
}));
}
};
let cache = global_cache();
let mut info = ModelInfo::from(model);
info.loaded = cache.is_loaded(model);
pgrx::JsonB(serde_json::to_value(info).unwrap_or_default())
}
/// Set the default embedding model
///
/// # Arguments
/// * `model_name` - Name of the model to set as default
///
/// # Returns
/// true if set successfully
///
/// # Example
/// ```sql
/// SELECT ruvector_set_default_model('bge-small');
/// ```
#[pg_extern]
pub fn ruvector_set_default_model(model_name: &str) -> bool {
let model = match EmbeddingModel::from_name(model_name) {
Some(m) => m,
None => {
pgrx::warning!("Unknown model: {}", model_name);
return false;
}
};
global_cache().set_default_model(model);
pgrx::info!("Default model set to '{}'", model.name());
true
}
/// Get the current default embedding model name
///
/// # Returns
/// Name of the default model
///
/// # Example
/// ```sql
/// SELECT ruvector_default_model();
/// ```
#[pg_extern]
pub fn ruvector_default_model() -> String {
global_cache().default_model().name().to_string()
}
/// Get embedding cache statistics
///
/// # Returns
/// JSON object with cache statistics
///
/// # Example
/// ```sql
/// SELECT ruvector_embedding_stats();
/// ```
#[pg_extern]
pub fn ruvector_embedding_stats() -> pgrx::JsonB {
let cache = global_cache();
let loaded_models = cache.loaded_models();
pgrx::JsonB(serde_json::json!({
"loaded_model_count": loaded_models.len(),
"loaded_models": loaded_models.iter().map(|m| m.name()).collect::<Vec<_>>(),
"estimated_memory_mb": cache.estimated_memory_usage() / (1024 * 1024),
"default_model": cache.default_model().name(),
"available_model_count": EmbeddingModel::all().len(),
}))
}
/// Get embedding dimensions for a model
///
/// # Arguments
/// * `model_name` - Name of the model
///
/// # Returns
/// Number of dimensions, or -1 if model unknown
///
/// # Example
/// ```sql
/// SELECT ruvector_embedding_dims('all-MiniLM-L6-v2'); -- Returns 384
/// ```
#[pg_extern(immutable, parallel_safe)]
pub fn ruvector_embedding_dims(model_name: &str) -> i32 {
match EmbeddingModel::from_name(model_name) {
Some(m) => m.dimensions() as i32,
None => -1,
}
}
// ============================================================================
// Tests
// ============================================================================
#[cfg(feature = "pg_test")]
#[pg_schema]
mod tests {
use super::*;
#[pg_test]
fn test_embedding_models_list() {
let models: Vec<_> = ruvector_embedding_models().collect();
assert!(!models.is_empty());
}
#[pg_test]
fn test_model_info() {
let info = ruvector_model_info("all-MiniLM-L6-v2");
let json = info.0;
assert!(json.get("name").is_some());
assert!(json.get("dimensions").is_some());
}
#[pg_test]
fn test_default_model() {
let name = ruvector_default_model();
assert!(!name.is_empty());
}
#[pg_test]
fn test_embedding_dims() {
assert_eq!(ruvector_embedding_dims("all-MiniLM-L6-v2"), 384);
assert_eq!(ruvector_embedding_dims("bge-base"), 768);
assert_eq!(ruvector_embedding_dims("unknown"), -1);
}
}
@@ -0,0 +1,47 @@
//! Local embedding generation module for ruvector-postgres
//!
//! Provides text-to-embedding functionality using fastembed-rs (ONNX-based models).
//! Supports multiple embedding models with lazy loading and thread-safe caching.
//!
//! # Features
//!
//! - Local embedding generation (no external API calls)
//! - Multiple model support (MiniLM, BGE, MPNet)
//! - Lazy model loading (loads on first use)
//! - Thread-safe model cache
//! - Batch embedding for efficiency
//!
//! # SQL Functions
//!
//! ```sql
//! -- Generate embedding from text
//! SELECT ruvector_embed('Hello world');
//! SELECT ruvector_embed('Hello world', 'all-MiniLM-L6-v2');
//!
//! -- Batch embedding
//! SELECT ruvector_embed_batch(ARRAY['text1', 'text2']);
//!
//! -- List available models
//! SELECT * FROM ruvector_embedding_models();
//!
//! -- Model management
//! SELECT ruvector_load_model('all-MiniLM-L6-v2');
//! SELECT ruvector_model_info('all-MiniLM-L6-v2');
//! ```
mod cache;
mod functions;
mod models;
pub use cache::ModelCache;
pub use functions::*;
pub use models::{EmbeddingModel, ModelInfo};
/// Default embedding model
pub const DEFAULT_MODEL: &str = "all-MiniLM-L6-v2";
/// Maximum batch size for embedding generation
pub const MAX_BATCH_SIZE: usize = 256;
/// Maximum text length (in characters) for embedding
pub const MAX_TEXT_LENGTH: usize = 8192;
@@ -0,0 +1,207 @@
//! Embedding model definitions and metadata
use serde::{Deserialize, Serialize};
/// Supported embedding models
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum EmbeddingModel {
/// all-MiniLM-L6-v2: Fast, good quality, 384 dimensions
AllMiniLmL6V2,
/// BAAI/bge-small-en-v1.5: Fast, high quality, 384 dimensions
BgeSmallEnV15,
/// BAAI/bge-base-en-v1.5: Medium speed, higher quality, 768 dimensions
BgeBaseEnV15,
/// sentence-transformers/all-mpnet-base-v2: Medium speed, high quality, 768 dimensions
AllMpnetBaseV2,
/// nomic-ai/nomic-embed-text-v1.5: Good quality, 768 dimensions
NomicEmbedTextV15,
/// BAAI/bge-large-en-v1.5: Slower, highest quality, 1024 dimensions
BgeLargeEnV15,
}
impl EmbeddingModel {
/// Parse model name string to enum
pub fn from_name(name: &str) -> Option<Self> {
match name.to_lowercase().as_str() {
"all-minilm-l6-v2" | "minilm" | "default" => Some(Self::AllMiniLmL6V2),
"bge-small-en-v1.5" | "bge-small" | "baai/bge-small-en-v1.5" => {
Some(Self::BgeSmallEnV15)
}
"bge-base-en-v1.5" | "bge-base" | "baai/bge-base-en-v1.5" => Some(Self::BgeBaseEnV15),
"bge-large-en-v1.5" | "bge-large" | "baai/bge-large-en-v1.5" => {
Some(Self::BgeLargeEnV15)
}
"all-mpnet-base-v2" | "mpnet" | "sentence-transformers/all-mpnet-base-v2" => {
Some(Self::AllMpnetBaseV2)
}
"nomic-embed-text-v1.5" | "nomic" | "nomic-ai/nomic-embed-text-v1.5" => {
Some(Self::NomicEmbedTextV15)
}
_ => None,
}
}
/// Get the canonical name for this model
pub fn name(&self) -> &'static str {
match self {
Self::AllMiniLmL6V2 => "all-MiniLM-L6-v2",
Self::BgeSmallEnV15 => "BAAI/bge-small-en-v1.5",
Self::BgeBaseEnV15 => "BAAI/bge-base-en-v1.5",
Self::BgeLargeEnV15 => "BAAI/bge-large-en-v1.5",
Self::AllMpnetBaseV2 => "sentence-transformers/all-mpnet-base-v2",
Self::NomicEmbedTextV15 => "nomic-ai/nomic-embed-text-v1.5",
}
}
/// Get the embedding dimensions for this model
pub fn dimensions(&self) -> usize {
match self {
Self::AllMiniLmL6V2 => 384,
Self::BgeSmallEnV15 => 384,
Self::BgeBaseEnV15 => 768,
Self::BgeLargeEnV15 => 1024,
Self::AllMpnetBaseV2 => 768,
Self::NomicEmbedTextV15 => 768,
}
}
/// Get a description of this model
pub fn description(&self) -> &'static str {
match self {
Self::AllMiniLmL6V2 => "Fast general-purpose model, good for most use cases",
Self::BgeSmallEnV15 => "High quality small model from BAAI, great for semantic search",
Self::BgeBaseEnV15 => "Higher quality base model from BAAI, better accuracy",
Self::BgeLargeEnV15 => "Highest quality large model from BAAI, best accuracy",
Self::AllMpnetBaseV2 => "High quality model from sentence-transformers",
Self::NomicEmbedTextV15 => "Modern model with good quality from Nomic AI",
}
}
/// Get model speed rating (1-5, higher is faster)
pub fn speed_rating(&self) -> u8 {
match self {
Self::AllMiniLmL6V2 => 5,
Self::BgeSmallEnV15 => 5,
Self::BgeBaseEnV15 => 3,
Self::BgeLargeEnV15 => 1,
Self::AllMpnetBaseV2 => 3,
Self::NomicEmbedTextV15 => 3,
}
}
/// Get model quality rating (1-5, higher is better)
pub fn quality_rating(&self) -> u8 {
match self {
Self::AllMiniLmL6V2 => 3,
Self::BgeSmallEnV15 => 4,
Self::BgeBaseEnV15 => 4,
Self::BgeLargeEnV15 => 5,
Self::AllMpnetBaseV2 => 4,
Self::NomicEmbedTextV15 => 4,
}
}
/// Get approximate memory usage in MB
pub fn memory_mb(&self) -> usize {
match self {
Self::AllMiniLmL6V2 => 90,
Self::BgeSmallEnV15 => 130,
Self::BgeBaseEnV15 => 440,
Self::BgeLargeEnV15 => 1340,
Self::AllMpnetBaseV2 => 440,
Self::NomicEmbedTextV15 => 550,
}
}
/// Get all supported models
pub fn all() -> &'static [EmbeddingModel] {
&[
Self::AllMiniLmL6V2,
Self::BgeSmallEnV15,
Self::BgeBaseEnV15,
Self::BgeLargeEnV15,
Self::AllMpnetBaseV2,
Self::NomicEmbedTextV15,
]
}
/// Get the default model
pub fn default_model() -> Self {
Self::AllMiniLmL6V2
}
}
impl Default for EmbeddingModel {
fn default() -> Self {
Self::default_model()
}
}
/// Model information for SQL queries
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelInfo {
pub name: String,
pub dimensions: i32,
pub description: String,
pub speed_rating: i32,
pub quality_rating: i32,
pub memory_mb: i32,
pub loaded: bool,
}
impl From<EmbeddingModel> for ModelInfo {
fn from(model: EmbeddingModel) -> Self {
Self {
name: model.name().to_string(),
dimensions: model.dimensions() as i32,
description: model.description().to_string(),
speed_rating: model.speed_rating() as i32,
quality_rating: model.quality_rating() as i32,
memory_mb: model.memory_mb() as i32,
loaded: false,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_model_parsing() {
assert_eq!(
EmbeddingModel::from_name("all-minilm-l6-v2"),
Some(EmbeddingModel::AllMiniLmL6V2)
);
assert_eq!(
EmbeddingModel::from_name("minilm"),
Some(EmbeddingModel::AllMiniLmL6V2)
);
assert_eq!(
EmbeddingModel::from_name("default"),
Some(EmbeddingModel::AllMiniLmL6V2)
);
assert_eq!(
EmbeddingModel::from_name("bge-small"),
Some(EmbeddingModel::BgeSmallEnV15)
);
assert_eq!(EmbeddingModel::from_name("unknown"), None);
}
#[test]
fn test_model_dimensions() {
assert_eq!(EmbeddingModel::AllMiniLmL6V2.dimensions(), 384);
assert_eq!(EmbeddingModel::BgeBaseEnV15.dimensions(), 768);
assert_eq!(EmbeddingModel::BgeLargeEnV15.dimensions(), 1024);
}
#[test]
fn test_all_models() {
let models = EmbeddingModel::all();
assert!(models.len() >= 4);
for model in models {
assert!(!model.name().is_empty());
assert!(model.dimensions() > 0);
}
}
}