use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use async_trait::async_trait;
use tokio::time::{Duration, sleep};
use xz_embed::{
ConcurrentBatchManager, EmbedError, EmbedModelInfo, EmbedPricing, EmbeddingModel, MockEmbedder,
RetryConfig,
};
#[derive(Debug)]
struct FailThenSucceedEmbedder {
info: EmbedModelInfo,
remaining_failures: AtomicU32,
}
impl FailThenSucceedEmbedder {
fn new(dimensions: usize, max_batch_size: usize, fail_count: u32) -> Self {
Self {
info: EmbedModelInfo {
name: "fail-then-succeed".into(),
display_name: "Fail Then Succeed".into(),
supported_dimensions: None,
current_dimension: dimensions,
max_input_tokens: 1024,
max_batch_size,
pricing: EmbedPricing { input_per_million: 0.0 },
},
remaining_failures: AtomicU32::new(fail_count),
}
}
}
#[async_trait]
impl EmbeddingModel for FailThenSucceedEmbedder {
async fn embed(&self, input: &[&str]) -> Result<Vec<Vec<f32>>, EmbedError> {
loop {
let current = self.remaining_failures.load(Ordering::Acquire);
if current == 0 {
break; }
if self
.remaining_failures
.compare_exchange(current, current - 1, Ordering::Release, Ordering::Acquire)
.is_ok()
{
return Err(EmbedError::Network("simulated transient failure".into()));
}
}
Ok(vec![vec![0.0; self.info.current_dimension]; input.len()])
}
fn model_info(&self) -> &EmbedModelInfo {
&self.info
}
fn max_batch_size(&self) -> usize {
self.info.max_batch_size
}
fn dimensions(&self) -> usize {
self.info.current_dimension
}
}
#[derive(Debug)]
struct ConcurrencyTrackerEmbedder {
info: EmbedModelInfo,
current_concurrent: Arc<AtomicU32>,
max_concurrent: Arc<AtomicU32>,
}
impl ConcurrencyTrackerEmbedder {
fn new(
dimensions: usize,
max_batch_size: usize,
current: Arc<AtomicU32>,
max: Arc<AtomicU32>,
) -> Self {
Self {
info: EmbedModelInfo {
name: "concurrency-tracker".into(),
display_name: "Concurrency Tracker".into(),
supported_dimensions: None,
current_dimension: dimensions,
max_input_tokens: 1024,
max_batch_size,
pricing: EmbedPricing { input_per_million: 0.0 },
},
current_concurrent: current,
max_concurrent: max,
}
}
}
#[async_trait]
impl EmbeddingModel for ConcurrencyTrackerEmbedder {
async fn embed(&self, input: &[&str]) -> Result<Vec<Vec<f32>>, EmbedError> {
let c = self.current_concurrent.fetch_add(1, Ordering::AcqRel) + 1;
let mut prev = self.max_concurrent.load(Ordering::Acquire);
while c > prev {
match self.max_concurrent.compare_exchange(
prev,
c,
Ordering::Release,
Ordering::Acquire,
) {
Ok(_) => break,
Err(actual) => prev = actual,
}
}
sleep(Duration::from_millis(50)).await;
self.current_concurrent.fetch_sub(1, Ordering::AcqRel);
Ok(vec![vec![0.0; self.info.current_dimension]; input.len()])
}
fn model_info(&self) -> &EmbedModelInfo {
&self.info
}
fn max_batch_size(&self) -> usize {
self.info.max_batch_size
}
fn dimensions(&self) -> usize {
self.info.current_dimension
}
}
#[tokio::test]
async fn test_batch_success() {
let mock = MockEmbedder::new(4, 32);
let manager = ConcurrentBatchManager::new(Box::new(mock), 2, 2);
let texts = vec!["a", "b", "c", "d"];
let results = manager.embed_all(&texts).await.unwrap();
assert_eq!(results.len(), 4, "should return 4 vectors");
for v in &results {
assert_eq!(v.len(), 4, "each vector should have dimension 4");
}
}
#[tokio::test]
async fn test_batch_retry() {
let embedder = FailThenSucceedEmbedder::new(4, 32, 2); let manager = ConcurrentBatchManager::new(Box::new(embedder), 10, 4).with_retry(RetryConfig {
max_retries: 3,
initial_backoff_ms: 1,
max_backoff_ms: 10,
backoff_multiplier: 2.0,
});
let texts = vec!["hello", "world"];
let result = manager.embed_all(&texts).await;
assert!(result.is_ok(), "batch should succeed after retries");
let vectors = result.unwrap();
assert_eq!(vectors.len(), 2, "should return 2 vectors");
assert_eq!(vectors[0].len(), 4);
assert_eq!(vectors[1].len(), 4);
}
#[tokio::test]
async fn test_batch_semaphore_limit() {
let current = Arc::new(AtomicU32::new(0));
let max = Arc::new(AtomicU32::new(0));
let embedder = ConcurrencyTrackerEmbedder::new(4, 2, current.clone(), max.clone());
let manager = ConcurrentBatchManager::new(Box::new(embedder), 2, 2);
let texts = vec!["a", "b", "c", "d", "e"];
let result = manager.embed_all(&texts).await;
assert!(result.is_ok(), "batch should complete successfully");
assert_eq!(result.unwrap().len(), 5, "should return 5 vectors");
let observed_max = max.load(Ordering::Acquire);
assert!(observed_max <= 2, "max concurrency observed was {observed_max}, expected ≤ 2");
}