xz-embed 0.1.1

文本向量嵌入与向量存储抽象层
Documentation
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,
};

// ═══════════════════════════════════════════════════════════════
// Test helper: an embedder that fails the first N calls then succeeds
// ═══════════════════════════════════════════════════════════════

#[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> {
        // Check-and-decrement: if there are remaining failures, fail with a retryable error
        loop {
            let current = self.remaining_failures.load(Ordering::Acquire);
            if current == 0 {
                break; // no more failures — proceed to succeed
            }
            if self
                .remaining_failures
                .compare_exchange(current, current - 1, Ordering::Release, Ordering::Acquire)
                .is_ok()
            {
                return Err(EmbedError::Network("simulated transient failure".into()));
            }
            // CAS failed — another thread decremented, retry
        }

        // Succeed: return zero vectors
        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
    }
}

// ═══════════════════════════════════════════════════════════════
// Test helper: an embedder that tracks concurrent calls
// ═══════════════════════════════════════════════════════════════

#[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;
        // Update observed max
        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 long enough for other concurrent tasks to overlap
        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
    }
}

// ═══════════════════════════════════════════════════════════════
// ConcurrentBatchManager tests
// ═══════════════════════════════════════════════════════════════

/// Basic concurrent batch processing succeeds and returns results
/// in the same order as the input texts.
#[tokio::test]
async fn test_batch_success() {
    // Use default MockEmbedder (no set_output) → returns zero vectors of
    // configured dimension for each per-batch call.
    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");
    }
}

/// Simulate transient failures and verify the manager retries
/// before eventually succeeding.
#[tokio::test]
async fn test_batch_retry() {
    let embedder = FailThenSucceedEmbedder::new(4, 32, 2); // fail twice, then succeed
    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);
}

/// Verify that the semaphore limits concurrent executions to
/// `max_concurrency` (here: 2).
#[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);

    // 5 texts → 3 batches (batch_size=2, so batches: [a,b], [c,d], [e])
    // With max_concurrency=2, at most 2 batches execute simultaneously.
    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");
}