xz-embed 0.1.1

文本向量嵌入与向量存储抽象层
Documentation
use wiremock::matchers::{body_json, header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use xz_embed::{
    EmbedConfig, EmbedError, EmbeddingModel, IndexBuildConfig, IndexBuildMode, ModelConfig,
    OpenAiEmbedder, RebuildTrigger, StorageConfig, TruncationStrategy,
};

/// Create an EmbedConfig pointing at a given base_url.
fn make_config(base_url: &str) -> EmbedConfig {
    EmbedConfig {
        default_model: "text-embedding-3-small".into(),
        default_dimensions: 4,
        truncation: TruncationStrategy::Error,
        storage: StorageConfig {
            backend: "memory".into(),
            path: String::new(),
            max_capacity_bytes: None,
            table_name: "embeddings".into(),
        },
        index: IndexBuildConfig {
            mode: IndexBuildMode::Manual,
            rebuild_trigger: RebuildTrigger::ManualOnly,
        },
        models: vec![ModelConfig {
            provider: "text-embedding-3-small".into(),
            api_key: Some("test-key".into()),
            model: Some("text-embedding-3-small".into()),
            dimensions: 4,
            max_batch_size: 2048,
            max_input_tokens: 8191,
            input_per_million: 0.02,
            model_path: None,
            tokenizer_path: None,
            base_url: Some(base_url.to_string()),
        }],
    }
}

/// Mock a successful POST /embeddings and verify the embedder returns parsed vectors.
#[tokio::test]
async fn test_openai_embed_success() {
    let mock_server = MockServer::start().await;

    Mock::given(method("POST"))
        .and(path("/embeddings"))
        .and(header("Authorization", "Bearer test-key"))
        .and(header("Content-Type", "application/json"))
        .and(body_json(serde_json::json!({
            "model": "text-embedding-3-small",
            "input": ["hello world", "test text"],
            "dimensions": 4,
        })))
        .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
            "data": [
                {"embedding": [0.1, 0.2, 0.3, 0.4], "index": 0},
                {"embedding": [0.5, 0.6, 0.7, 0.8], "index": 1}
            ],
            "model": "text-embedding-3-small",
            "usage": {"prompt_tokens": 4, "total_tokens": 4}
        })))
        .mount(&mock_server)
        .await;

    let config = make_config(&mock_server.uri());
    let embedder = OpenAiEmbedder::from_config(&config).unwrap();

    let vectors = embedder.embed(&["hello world", "test text"]).await.unwrap();
    assert_eq!(vectors.len(), 2);
    assert_eq!(vectors[0], vec![0.1, 0.2, 0.3, 0.4]);
    assert_eq!(vectors[1], vec![0.5, 0.6, 0.7, 0.8]);
}

/// Empty input should return EmbedError::EmptyBatch without any network call.
#[tokio::test]
async fn test_openai_embed_empty_batch() {
    let mock_server = MockServer::start().await;
    let config = make_config(&mock_server.uri());
    let embedder = OpenAiEmbedder::from_config(&config).unwrap();

    let result = embedder.embed(&[]).await;
    let err = result.unwrap_err();
    assert!(matches!(err, EmbedError::EmptyBatch));
}

/// 401 response should produce EmbedError::Auth.
#[tokio::test]
async fn test_openai_embed_auth_error() {
    let mock_server = MockServer::start().await;

    Mock::given(method("POST"))
        .and(path("/embeddings"))
        .respond_with(ResponseTemplate::new(401).set_body_string("Unauthorized"))
        .mount(&mock_server)
        .await;

    let config = make_config(&mock_server.uri());
    let embedder = OpenAiEmbedder::from_config(&config).unwrap();

    let result = embedder.embed(&["hello world"]).await;
    let err = result.unwrap_err();
    assert!(matches!(err, EmbedError::Auth(_)));
}

/// Response with wrong dimension count should produce EmbedError::DimensionMismatch.
#[tokio::test]
async fn test_openai_embed_dimension_mismatch() {
    let mock_server = MockServer::start().await;

    Mock::given(method("POST"))
        .and(path("/embeddings"))
        .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
            "data": [
                {"embedding": [0.1, 0.2, 0.3], "index": 0}
            ],
            "model": "text-embedding-3-small",
            "usage": {"prompt_tokens": 2, "total_tokens": 2}
        })))
        .mount(&mock_server)
        .await;

    let config = make_config(&mock_server.uri());
    let embedder = OpenAiEmbedder::from_config(&config).unwrap();

    let result = embedder.embed(&["hello world"]).await;
    let err = result.unwrap_err();
    assert!(
        matches!(err, EmbedError::DimensionMismatch { expected: 4, actual: 3 }),
        "expected DimensionMismatch(4, 3), got {err:?}",
    );
}

/// HTTP 429 rate-limit response should produce EmbedError::RateLimit.
#[tokio::test]
async fn test_openai_embed_rate_limit() {
    let mock_server = MockServer::start().await;

    Mock::given(method("POST"))
        .and(path("/embeddings"))
        .respond_with(
            ResponseTemplate::new(429)
                .insert_header("retry-after", "5000")
                .set_body_string("Too Many Requests"),
        )
        .mount(&mock_server)
        .await;

    let config = make_config(&mock_server.uri());
    let embedder = OpenAiEmbedder::from_config(&config).unwrap();

    let result = embedder.embed(&["hello world"]).await;
    let err = result.unwrap_err();
    assert!(
        matches!(err, EmbedError::RateLimit { retry_after_ms: 5000 }),
        "expected RateLimit(5000), got {err:?}",
    );
}