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,
};
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()),
}],
}
}
#[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]);
}
#[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));
}
#[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(_)));
}
#[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:?}",
);
}
#[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:?}",
);
}