use crate::error::AiError;
use async_trait::async_trait;
use parking_lot::RwLock;
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct EmbeddingError {
pub message: String,
pub model: Option<String>,
}
impl EmbeddingError {
pub fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
model: None,
}
}
pub fn with_model(message: impl Into<String>, model: impl Into<String>) -> Self {
Self {
message: message.into(),
model: Some(model.into()),
}
}
}
impl std::fmt::Display for EmbeddingError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "EmbeddingError: {}", self.message)?;
if let Some(ref model) = self.model {
write!(f, " (model: {})", model)?;
}
Ok(())
}
}
impl std::error::Error for EmbeddingError {}
#[async_trait]
pub trait EmbeddingModel: Send + Sync {
async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError>;
async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError>;
fn dimension(&self) -> usize;
fn model_name(&self) -> &str;
}
pub struct EmbeddingRecord {
pub id: String,
pub text: String,
pub vector: Vec<f32>,
pub metadata: Option<std::collections::HashMap<String, serde_json::Value>>,
}
impl EmbeddingRecord {
pub fn new(id: impl Into<String>, text: impl Into<String>, vector: Vec<f32>) -> Self {
Self {
id: id.into(),
text: text.into(),
vector,
metadata: None,
}
}
pub fn with_metadata(
id: impl Into<String>,
text: impl Into<String>,
vector: Vec<f32>,
metadata: std::collections::HashMap<String, serde_json::Value>,
) -> Self {
Self {
id: id.into(),
text: text.into(),
vector,
metadata: Some(metadata),
}
}
}
pub struct EmbeddingBatch {
pub records: Vec<EmbeddingRecord>,
pub batch_size: usize,
}
impl EmbeddingBatch {
pub fn new(records: Vec<EmbeddingRecord>) -> Self {
Self {
records,
batch_size: 32,
}
}
pub fn with_batch_size(mut self, size: usize) -> Self {
self.batch_size = size;
self
}
pub fn batch_chunks(&self) -> Vec<&[EmbeddingRecord]> {
self.records.chunks(self.batch_size).collect()
}
}
pub struct SimpleEmbeddingModel {
name: String,
dimension: usize,
vocabulary: RwLock<HashMap<String, usize>>,
}
impl SimpleEmbeddingModel {
pub fn new(name: impl Into<String>, dimension: usize) -> Self {
Self {
name: name.into(),
dimension,
vocabulary: RwLock::new(HashMap::new()),
}
}
pub fn vocabulary_size(&self) -> usize {
self.vocabulary.read().len()
}
fn register_token(&self, token: &str) -> usize {
let mut vocab = self.vocabulary.write();
if let Some(&idx) = vocab.get(token) {
return idx;
}
let idx = fnv1a(token) % self.dimension.max(1);
vocab.insert(token.to_string(), idx);
idx
}
fn tokenize(text: &str) -> Vec<String> {
text.split(|c: char| !c.is_alphanumeric())
.filter(|s| !s.is_empty())
.map(|s| s.to_lowercase())
.collect()
}
fn embed_text(&self, text: &str) -> Vec<f32> {
let mut vec = vec![0.0f32; self.dimension];
if self.dimension == 0 {
return vec;
}
let tokens = Self::tokenize(text);
if tokens.is_empty() {
return vec;
}
for token in &tokens {
let idx = self.register_token(token);
vec[idx] += 1.0;
}
let norm: f32 = vec.iter().map(|v| v * v).sum::<f32>().sqrt();
if norm > 0.0 {
for v in vec.iter_mut() {
*v /= norm;
}
}
vec
}
}
fn fnv1a(s: &str) -> usize {
let mut hash: u64 = 0xcbf29ce484222325;
for byte in s.as_bytes() {
hash ^= *byte as u64;
hash = hash.wrapping_mul(0x100000001b3);
}
hash as usize
}
pub struct CachingEmbeddingModel<M>
where
M: EmbeddingModel,
{
inner: M,
cache: RwLock<HashMap<String, Vec<f32>>>,
hits: std::sync::atomic::AtomicU64,
misses: std::sync::atomic::AtomicU64,
}
impl<M> CachingEmbeddingModel<M>
where
M: EmbeddingModel,
{
pub fn new(inner: M) -> Self {
Self {
inner,
cache: RwLock::new(HashMap::new()),
hits: std::sync::atomic::AtomicU64::new(0),
misses: std::sync::atomic::AtomicU64::new(0),
}
}
pub fn cache_size(&self) -> usize {
self.cache.read().len()
}
pub fn cache_hits(&self) -> u64 {
self.hits.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn cache_misses(&self) -> u64 {
self.misses.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn hit_rate(&self) -> f64 {
let hits = self.cache_hits();
let misses = self.cache_misses();
let total = hits + misses;
if total == 0 {
return 0.0;
}
hits as f64 / total as f64
}
pub fn clear_cache(&self) {
let mut cache = self.cache.write();
cache.clear();
}
pub fn inner(&self) -> &M {
&self.inner
}
}
#[async_trait]
impl<M> EmbeddingModel for CachingEmbeddingModel<M>
where
M: EmbeddingModel,
{
async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError> {
{
let cache = self.cache.read();
if let Some(vector) = cache.get(text) {
self.hits.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(vector.clone());
}
}
self.misses
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let vector = self.inner.embed(text).await?;
{
let mut cache = self.cache.write();
cache.insert(text.to_string(), vector.clone());
}
Ok(vector)
}
async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError> {
let mut results = Vec::with_capacity(texts.len());
let mut uncached_indices = Vec::new();
let mut uncached_texts = Vec::new();
{
let cache = self.cache.read();
for (idx, text) in texts.iter().enumerate() {
if let Some(vector) = cache.get(text) {
self.hits.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
results.push(vector.clone());
} else {
self.misses
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
uncached_indices.push(idx);
uncached_texts.push(text.clone());
results.push(Vec::new()); }
}
}
if !uncached_texts.is_empty() {
let vectors = self.inner.embed_batch(&uncached_texts).await?;
let mut cache = self.cache.write();
for (i, idx) in uncached_indices.iter().enumerate() {
let text = &uncached_texts[i];
let vector = &vectors[i];
results[*idx] = vector.clone();
cache.insert(text.clone(), vector.clone());
}
}
Ok(results)
}
fn dimension(&self) -> usize {
self.inner.dimension()
}
fn model_name(&self) -> &str {
self.inner.model_name()
}
}
pub struct NormalizedEmbeddingModel<M>
where
M: EmbeddingModel,
{
inner: M,
}
impl<M> NormalizedEmbeddingModel<M>
where
M: EmbeddingModel,
{
pub fn new(inner: M) -> Self {
Self { inner }
}
pub fn l2_normalize(vector: &mut [f32]) {
let norm: f32 = vector.iter().map(|v| v * v).sum::<f32>().sqrt();
if norm > 0.0 {
for v in vector.iter_mut() {
*v /= norm;
}
}
}
pub fn inner(&self) -> &M {
&self.inner
}
}
#[async_trait]
impl<M> EmbeddingModel for NormalizedEmbeddingModel<M>
where
M: EmbeddingModel,
{
async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError> {
let mut vector = self.inner.embed(text).await?;
Self::l2_normalize(&mut vector);
Ok(vector)
}
async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError> {
let mut vectors = self.inner.embed_batch(texts).await?;
for vector in vectors.iter_mut() {
Self::l2_normalize(vector);
}
Ok(vectors)
}
fn dimension(&self) -> usize {
self.inner.dimension()
}
fn model_name(&self) -> &str {
self.inner.model_name()
}
}
pub struct DimReductionEmbeddingModel<M>
where
M: EmbeddingModel,
{
inner: M,
target_dimension: usize,
strategy: DimReductionStrategy,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DimReductionStrategy {
Truncate,
Average,
}
impl<M> DimReductionEmbeddingModel<M>
where
M: EmbeddingModel,
{
pub fn new(inner: M, target_dimension: usize, strategy: DimReductionStrategy) -> Self {
Self {
inner,
target_dimension,
strategy,
}
}
pub fn truncate(inner: M, target_dimension: usize) -> Self {
Self::new(inner, target_dimension, DimReductionStrategy::Truncate)
}
pub fn average(inner: M, target_dimension: usize) -> Self {
Self::new(inner, target_dimension, DimReductionStrategy::Average)
}
pub fn reduce(&self, vector: Vec<f32>) -> Vec<f32> {
match self.strategy {
DimReductionStrategy::Truncate => {
vector.into_iter().take(self.target_dimension).collect()
}
DimReductionStrategy::Average => {
if vector.is_empty() || self.target_dimension == 0 {
return Vec::new();
}
let chunk_size = vector.len() / self.target_dimension;
if chunk_size == 0 {
return vector.into_iter().take(self.target_dimension).collect();
}
let mut result = Vec::with_capacity(self.target_dimension);
for i in 0..self.target_dimension {
let start = i * chunk_size;
let end = if i == self.target_dimension - 1 {
vector.len()
} else {
start + chunk_size
};
let chunk = &vector[start..end];
let avg: f32 = chunk.iter().sum::<f32>() / chunk.len() as f32;
result.push(avg);
}
result
}
}
}
pub fn inner(&self) -> &M {
&self.inner
}
pub fn target_dimension(&self) -> usize {
self.target_dimension
}
pub fn strategy(&self) -> DimReductionStrategy {
self.strategy
}
}
#[async_trait]
impl<M> EmbeddingModel for DimReductionEmbeddingModel<M>
where
M: EmbeddingModel,
{
async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError> {
let vector = self.inner.embed(text).await?;
Ok(self.reduce(vector))
}
async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError> {
let vectors = self.inner.embed_batch(texts).await?;
Ok(vectors.into_iter().map(|v| self.reduce(v)).collect())
}
fn dimension(&self) -> usize {
self.target_dimension
}
fn model_name(&self) -> &str {
self.inner.model_name()
}
}
pub struct LoggingEmbeddingModel<M>
where
M: EmbeddingModel,
{
inner: M,
call_count: std::sync::atomic::AtomicU64,
total_texts: std::sync::atomic::AtomicU64,
}
impl<M> LoggingEmbeddingModel<M>
where
M: EmbeddingModel,
{
pub fn new(inner: M) -> Self {
Self {
inner,
call_count: std::sync::atomic::AtomicU64::new(0),
total_texts: std::sync::atomic::AtomicU64::new(0),
}
}
pub fn call_count(&self) -> u64 {
self.call_count.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn total_texts(&self) -> u64 {
self.total_texts.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn inner(&self) -> &M {
&self.inner
}
}
#[async_trait]
impl<M> EmbeddingModel for LoggingEmbeddingModel<M>
where
M: EmbeddingModel,
{
async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError> {
self.call_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.total_texts
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.inner.embed(text).await
}
async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError> {
self.call_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.total_texts
.fetch_add(texts.len() as u64, std::sync::atomic::Ordering::Relaxed);
self.inner.embed_batch(texts).await
}
fn dimension(&self) -> usize {
self.inner.dimension()
}
fn model_name(&self) -> &str {
self.inner.model_name()
}
}
#[async_trait]
impl EmbeddingModel for SimpleEmbeddingModel {
async fn embed(&self, text: &str) -> Result<Vec<f32>, AiError> {
Ok(self.embed_text(text))
}
async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, AiError> {
Ok(texts.iter().map(|t| self.embed_text(t)).collect())
}
fn dimension(&self) -> usize {
self.dimension
}
fn model_name(&self) -> &str {
&self.name
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_embed_simple_text() {
let model = SimpleEmbeddingModel::new("test-model", 16);
let v = model.embed("hello world").await.unwrap();
assert_eq!(v.len(), 16);
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-5 || norm.abs() < 1e-5);
}
#[tokio::test]
async fn test_embed_empty_text() {
let model = SimpleEmbeddingModel::new("test-model", 8);
let v = model.embed("").await.unwrap();
assert!(v.iter().all(|x| *x == 0.0));
}
#[tokio::test]
async fn test_embed_deterministic() {
let model = SimpleEmbeddingModel::new("test-model", 32);
let v1 = model.embed("rust programming language").await.unwrap();
let v2 = model.embed("rust programming language").await.unwrap();
assert_eq!(v1, v2);
}
#[tokio::test]
async fn test_embed_similar_texts_closer_than_different() {
let model = SimpleEmbeddingModel::new("test-model", 64);
let v1 = model.embed("the quick brown fox jumps").await.unwrap();
let v2 = model.embed("the quick brown fox").await.unwrap();
let v3 = model
.embed("completely different words here")
.await
.unwrap();
let sim_close = cosine(&v1, &v2);
let sim_far = cosine(&v1, &v3);
assert!(
sim_close >= sim_far,
"similar texts should be at least as close"
);
}
#[tokio::test]
async fn test_embed_batch() {
let model = SimpleEmbeddingModel::new("test-model", 16);
let texts = vec!["hello".to_string(), "world".to_string()];
let vecs = model.embed_batch(&texts).await.unwrap();
assert_eq!(vecs.len(), 2);
assert_eq!(vecs[0].len(), 16);
assert_eq!(vecs[1].len(), 16);
}
#[test]
fn test_dimension_and_name() {
let model = SimpleEmbeddingModel::new("my-model", 128);
assert_eq!(model.dimension(), 128);
assert_eq!(model.model_name(), "my-model");
}
fn cosine(a: &[f32], b: &[f32]) -> f32 {
let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if na == 0.0 || nb == 0.0 {
return 0.0;
}
dot / (na * nb)
}
#[tokio::test]
async fn test_caching_model_caches_repeated_calls() {
let inner = SimpleEmbeddingModel::new("test", 16);
let caching = CachingEmbeddingModel::new(inner);
let v1 = caching.embed("hello world").await.unwrap();
let v2 = caching.embed("hello world").await.unwrap();
assert_eq!(v1, v2);
assert_eq!(caching.cache_hits(), 1);
assert_eq!(caching.cache_misses(), 1);
assert_eq!(caching.cache_size(), 1);
}
#[tokio::test]
async fn test_caching_model_different_texts() {
let inner = SimpleEmbeddingModel::new("test", 16);
let caching = CachingEmbeddingModel::new(inner);
let v1 = caching.embed("hello").await.unwrap();
let v2 = caching.embed("world").await.unwrap();
assert_eq!(v1.len(), 16);
assert_eq!(v2.len(), 16);
assert_ne!(v1, v2, "different inputs must yield different embeddings");
assert_eq!(caching.cache_misses(), 2);
assert_eq!(caching.cache_hits(), 0);
assert_eq!(caching.cache_size(), 2);
}
#[tokio::test]
async fn test_caching_model_hit_rate() {
let inner = SimpleEmbeddingModel::new("test", 16);
let caching = CachingEmbeddingModel::new(inner);
let v1_miss = caching.embed("a").await.unwrap();
let v1_hit = caching.embed("a").await.unwrap(); let v2_miss = caching.embed("b").await.unwrap();
let v1_hit2 = caching.embed("a").await.unwrap();
assert_eq!(v1_miss, v1_hit, "cache hit must return identical vector");
assert_eq!(v1_miss, v1_hit2, "cache hit must return identical vector");
assert_ne!(
v1_miss, v2_miss,
"different inputs must yield different vectors"
);
assert!((caching.hit_rate() - 0.5).abs() < 1e-6);
}
#[tokio::test]
async fn test_caching_model_hit_rate_zero_when_empty() {
let inner = SimpleEmbeddingModel::new("test", 16);
let caching = CachingEmbeddingModel::new(inner);
assert_eq!(caching.hit_rate(), 0.0);
}
#[tokio::test]
async fn test_caching_model_clear_cache() {
let inner = SimpleEmbeddingModel::new("test", 16);
let caching = CachingEmbeddingModel::new(inner);
let v1 = caching.embed("hello").await.unwrap();
assert_eq!(caching.cache_size(), 1);
caching.clear_cache();
assert_eq!(caching.cache_size(), 0);
let v2 = caching.embed("hello").await.unwrap();
assert_eq!(
v1, v2,
"embeddings must be deterministic across cache clears"
);
assert_eq!(caching.cache_misses(), 2);
}
#[tokio::test]
async fn test_caching_model_batch_mixed() {
let inner = SimpleEmbeddingModel::new("test", 16);
let caching = CachingEmbeddingModel::new(inner);
let seed = caching.embed("hello").await.unwrap();
let texts = vec!["hello".to_string(), "world".to_string()];
let results = caching.embed_batch(&texts).await.unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0], seed, "cached entry must match seed vector");
assert_eq!(results[1].len(), 16);
assert_ne!(results[0], results[1], "different inputs must differ");
assert_eq!(caching.cache_hits(), 1); assert_eq!(caching.cache_misses(), 2); }
#[tokio::test]
async fn test_caching_model_batch_all_cached() {
let inner = SimpleEmbeddingModel::new("test", 16);
let caching = CachingEmbeddingModel::new(inner);
let seed_hello = caching.embed("hello").await.unwrap();
let seed_world = caching.embed("world").await.unwrap();
let texts = vec!["hello".to_string(), "world".to_string()];
let results = caching.embed_batch(&texts).await.unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0], seed_hello, "cached hello must match seed");
assert_eq!(results[1], seed_world, "cached world must match seed");
assert_eq!(caching.cache_hits(), 2);
}
#[tokio::test]
async fn test_caching_model_preserves_dimension_and_name() {
let inner = SimpleEmbeddingModel::new("my-model", 32);
let caching = CachingEmbeddingModel::new(inner);
assert_eq!(caching.dimension(), 32);
assert_eq!(caching.model_name(), "my-model");
}
#[tokio::test]
async fn test_caching_model_inner_access() {
let inner = SimpleEmbeddingModel::new("inner", 8);
let caching = CachingEmbeddingModel::new(inner);
assert_eq!(caching.inner().model_name(), "inner");
}
#[tokio::test]
async fn test_normalized_model_produces_unit_vector() {
let inner = SimpleEmbeddingModel::new("test", 16);
let normalized = NormalizedEmbeddingModel::new(inner);
let v = normalized.embed("hello world").await.unwrap();
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(norm.abs() < 1e-5 || (norm - 1.0).abs() < 1e-5);
}
#[tokio::test]
async fn test_normalized_model_batch_produces_unit_vectors() {
let inner = SimpleEmbeddingModel::new("test", 16);
let normalized = NormalizedEmbeddingModel::new(inner);
let texts = vec!["hello".to_string(), "world".to_string()];
let vectors = normalized.embed_batch(&texts).await.unwrap();
for v in &vectors {
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(norm.abs() < 1e-5 || (norm - 1.0).abs() < 1e-5);
}
}
#[tokio::test]
async fn test_normalized_model_l2_normalize_static() {
let mut vector = vec![3.0, 4.0]; NormalizedEmbeddingModel::<SimpleEmbeddingModel>::l2_normalize(&mut vector);
assert!((vector[0] - 0.6).abs() < 1e-5);
assert!((vector[1] - 0.8).abs() < 1e-5);
}
#[tokio::test]
async fn test_normalized_model_l2_normalize_zero_vector() {
let mut vector = vec![0.0, 0.0, 0.0];
NormalizedEmbeddingModel::<SimpleEmbeddingModel>::l2_normalize(&mut vector);
assert!(vector.iter().all(|x| *x == 0.0));
}
#[tokio::test]
async fn test_normalized_model_preserves_dimension_and_name() {
let inner = SimpleEmbeddingModel::new("norm-model", 64);
let normalized = NormalizedEmbeddingModel::new(inner);
assert_eq!(normalized.dimension(), 64);
assert_eq!(normalized.model_name(), "norm-model");
}
#[tokio::test]
async fn test_normalized_model_inner_access() {
let inner = SimpleEmbeddingModel::new("inner", 8);
let normalized = NormalizedEmbeddingModel::new(inner);
assert_eq!(normalized.inner().model_name(), "inner");
}
#[test]
fn test_dim_reduction_strategy_variants() {
assert_eq!(
DimReductionStrategy::Truncate,
DimReductionStrategy::Truncate
);
assert_eq!(DimReductionStrategy::Average, DimReductionStrategy::Average);
assert_ne!(
DimReductionStrategy::Truncate,
DimReductionStrategy::Average
);
}
#[tokio::test]
async fn test_dim_reduction_truncate() {
let inner = SimpleEmbeddingModel::new("test", 16);
let reduced = DimReductionEmbeddingModel::truncate(inner, 8);
let v = reduced.embed("hello world").await.unwrap();
assert_eq!(v.len(), 8);
}
#[tokio::test]
async fn test_dim_reduction_truncate_batch() {
let inner = SimpleEmbeddingModel::new("test", 16);
let reduced = DimReductionEmbeddingModel::truncate(inner, 4);
let texts = vec!["hello".to_string(), "world".to_string()];
let vectors = reduced.embed_batch(&texts).await.unwrap();
for v in &vectors {
assert_eq!(v.len(), 4);
}
}
#[tokio::test]
async fn test_dim_reduction_average() {
let inner = SimpleEmbeddingModel::new("test", 16);
let reduced = DimReductionEmbeddingModel::average(inner, 4);
let v = reduced.embed("hello world").await.unwrap();
assert_eq!(v.len(), 4);
}
#[test]
fn test_dim_reduction_reduce_truncate() {
let inner = SimpleEmbeddingModel::new("test", 8);
let reduced = DimReductionEmbeddingModel::truncate(inner, 3);
let result = reduced.reduce(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
assert_eq!(result, vec![1.0, 2.0, 3.0]);
}
#[test]
fn test_dim_reduction_reduce_average() {
let inner = SimpleEmbeddingModel::new("test", 8);
let reduced = DimReductionEmbeddingModel::average(inner, 2);
let result = reduced.reduce(vec![1.0, 2.0, 3.0, 4.0]);
assert!((result[0] - 1.5).abs() < 1e-5);
assert!((result[1] - 3.5).abs() < 1e-5);
}
#[test]
fn test_dim_reduction_reduce_empty() {
let inner = SimpleEmbeddingModel::new("test", 8);
let reduced = DimReductionEmbeddingModel::average(inner, 2);
let result = reduced.reduce(vec![]);
assert!(result.is_empty());
}
#[test]
fn test_dim_reduction_reduce_target_zero() {
let inner = SimpleEmbeddingModel::new("test", 8);
let reduced = DimReductionEmbeddingModel::average(inner, 0);
let result = reduced.reduce(vec![1.0, 2.0, 3.0]);
assert!(result.is_empty());
}
#[test]
fn test_dim_reduction_reduce_truncate_smaller_than_target() {
let inner = SimpleEmbeddingModel::new("test", 8);
let reduced = DimReductionEmbeddingModel::truncate(inner, 10);
let result = reduced.reduce(vec![1.0, 2.0, 3.0]);
assert_eq!(result.len(), 3);
}
#[tokio::test]
async fn test_dim_reduction_dimension_returns_target() {
let inner = SimpleEmbeddingModel::new("test", 16);
let reduced = DimReductionEmbeddingModel::truncate(inner, 8);
assert_eq!(reduced.dimension(), 8);
}
#[tokio::test]
async fn test_dim_reduction_preserves_model_name() {
let inner = SimpleEmbeddingModel::new("original", 16);
let reduced = DimReductionEmbeddingModel::truncate(inner, 8);
assert_eq!(reduced.model_name(), "original");
}
#[test]
fn test_dim_reduction_target_dimension_and_strategy_accessors() {
let inner = SimpleEmbeddingModel::new("test", 16);
let reduced = DimReductionEmbeddingModel::new(inner, 8, DimReductionStrategy::Average);
assert_eq!(reduced.target_dimension(), 8);
assert_eq!(reduced.strategy(), DimReductionStrategy::Average);
}
#[tokio::test]
async fn test_logging_model_counts_calls() {
let inner = SimpleEmbeddingModel::new("test", 16);
let logging = LoggingEmbeddingModel::new(inner);
let v1 = logging.embed("hello").await.unwrap();
let v2 = logging.embed("world").await.unwrap();
let baseline = SimpleEmbeddingModel::new("test", 16);
assert_eq!(v1, baseline.embed_text("hello"));
assert_eq!(v2, baseline.embed_text("world"));
assert_ne!(v1, v2);
assert_eq!(logging.call_count(), 2);
assert_eq!(logging.total_texts(), 2);
}
#[tokio::test]
async fn test_logging_model_batch_counts() {
let inner = SimpleEmbeddingModel::new("test", 16);
let logging = LoggingEmbeddingModel::new(inner);
let texts = vec!["hello".to_string(), "world".to_string(), "foo".to_string()];
let results = logging.embed_batch(&texts).await.unwrap();
let baseline = SimpleEmbeddingModel::new("test", 16);
assert_eq!(results.len(), 3);
assert_eq!(results[0], baseline.embed_text("hello"));
assert_eq!(results[1], baseline.embed_text("world"));
assert_eq!(results[2], baseline.embed_text("foo"));
assert_eq!(logging.call_count(), 1);
assert_eq!(logging.total_texts(), 3);
}
#[tokio::test]
async fn test_logging_model_preserves_dimension_and_name() {
let inner = SimpleEmbeddingModel::new("logged", 32);
let logging = LoggingEmbeddingModel::new(inner);
assert_eq!(logging.dimension(), 32);
assert_eq!(logging.model_name(), "logged");
}
#[tokio::test]
async fn test_logging_model_inner_access() {
let inner = SimpleEmbeddingModel::new("inner", 8);
let logging = LoggingEmbeddingModel::new(inner);
assert_eq!(logging.inner().model_name(), "inner");
}
#[tokio::test]
async fn test_logging_model_initial_counts_zero() {
let inner = SimpleEmbeddingModel::new("test", 16);
let logging = LoggingEmbeddingModel::new(inner);
assert_eq!(logging.call_count(), 0);
assert_eq!(logging.total_texts(), 0);
}
#[tokio::test]
async fn test_compose_caching_and_normalized() {
let inner = SimpleEmbeddingModel::new("composed", 16);
let caching = CachingEmbeddingModel::new(inner);
let normalized = NormalizedEmbeddingModel::new(caching);
let v1 = normalized.embed("hello world").await.unwrap();
let v2 = normalized.embed("hello world").await.unwrap();
assert_eq!(v1, v2);
let norm: f32 = v1.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(norm.abs() < 1e-5 || (norm - 1.0).abs() < 1e-5);
}
#[tokio::test]
async fn test_compose_logging_and_caching() {
let inner = SimpleEmbeddingModel::new("composed", 16);
let logging = LoggingEmbeddingModel::new(inner);
let caching = CachingEmbeddingModel::new(logging);
let v1 = caching.embed("hello").await.unwrap();
let v2 = caching.embed("hello").await.unwrap();
assert_eq!(v1, v2, "cache hit must return identical vector");
let baseline = SimpleEmbeddingModel::new("composed", 16);
assert_eq!(v1, baseline.embed_text("hello"));
assert_eq!(caching.inner().call_count(), 1);
assert_eq!(caching.cache_hits(), 1);
}
}