From f19bd6e8005410eeb055463ca6d700872a731322 Mon Sep 17 00:00:00 2001 From: Brummel Date: Mon, 18 May 2026 18:46:32 +0200 Subject: [PATCH] fix: embed_batch length+index pairing; reject non-numeric components Co-Authored-By: Claude Sonnet 4.6 --- src/embed.rs | 18 +++++++- tests/embed_tests.rs | 103 ++++++++++++++++++++++++++++++++++++++++++- 2 files changed, 118 insertions(+), 3 deletions(-) diff --git a/src/embed.rs b/src/embed.rs index c80f7e3..2c65a7c 100644 --- a/src/embed.rs +++ b/src/embed.rs @@ -72,12 +72,26 @@ pub fn embed_batch( let resp: serde_json::Value = client.post_json("/embeddings", &body)?; let data = resp["data"].as_array() .ok_or_else(|| AppError::Ionos("embeddings: no data[]".into()))?; + if data.len() != texts.len() { + return Err(AppError::Ionos(format!( + "embeddings: expected {} items, got {}", texts.len(), data.len() + ))); + } let mut n = 0; for (i, item) in data.iter().enumerate() { + let idx = item["index"].as_u64().map(|u| u as usize).unwrap_or(i); + if idx >= texts.len() { + return Err(AppError::Ionos(format!( + "embeddings: index {idx} out of range (n={})", texts.len() + ))); + } let v: Vec = item["embedding"].as_array() .ok_or_else(|| AppError::Ionos("embeddings: no embedding[]".into()))? - .iter().map(|x| x.as_f64().unwrap_or(0.0) as f32).collect(); - store.put(&texts[i], &v)?; + .iter() + .map(|x| x.as_f64().ok_or_else(|| AppError::Ionos("embeddings: non-numeric component".into()))) + .collect::, _>>()? + .into_iter().map(|f| f as f32).collect(); + store.put(&texts[idx], &v)?; n += 1; } Ok(n) diff --git a/tests/embed_tests.rs b/tests/embed_tests.rs index 43bbf80..1071d5e 100644 --- a/tests/embed_tests.rs +++ b/tests/embed_tests.rs @@ -1,4 +1,5 @@ -use alpha_id::embed::{EmbeddingStore, estimate}; +use alpha_id::embed::{EmbeddingStore, embed_batch, estimate}; +use alpha_id::ionos::IonosClient; #[test] fn cache_key_is_stable_sha256() { @@ -26,3 +27,103 @@ fn estimate_reports_pending_count_and_chars() { assert_eq!(est.items, 2); assert_eq!(est.chars, 5); } + +// --------------------------------------------------------------------------- +// Regression tests for embed_batch hardening (Fix 1 & Fix 2) +// --------------------------------------------------------------------------- + +/// Spin up a tiny_http mock server that serves one fixed JSON response, then +/// return an IonosClient pointed at it together with a temp token file. +fn make_mock_client( + response_body: &str, + token_dir: &tempfile::TempDir, +) -> (IonosClient, std::net::SocketAddr, std::thread::JoinHandle<()>) { + use tiny_http::{Server, Response}; + + let server = Server::http("127.0.0.1:0").unwrap(); + let addr = server.server_addr().to_ip().unwrap(); + + // Write a dummy token file + let token_path = token_dir.path().join("token.txt"); + std::fs::write(&token_path, "dummy-token").unwrap(); + + let body = response_body.to_string(); + let handle = std::thread::spawn(move || { + if let Ok(req) = server.recv() { + let response = Response::from_string(body) + .with_header( + "Content-Type: application/json".parse::().unwrap(), + ); + let _ = req.respond(response); + } + }); + + let base_url = format!("http://{}", addr); + let client = IonosClient::new(&base_url, token_path.to_str().unwrap()).unwrap(); + (client, addr, handle) +} + +#[test] +fn embed_batch_happy_path_index_pairing() { + let token_dir = tempfile::tempdir().unwrap(); + let embed_dir = tempfile::tempdir().unwrap(); + + let body = r#"{"data":[{"index":0,"embedding":[0.1,0.2]},{"index":1,"embedding":[0.3,0.4]}]}"#; + let (client, _addr, handle) = make_mock_client(body, &token_dir); + + let store = EmbeddingStore::open(embed_dir.path().to_str().unwrap(), "test-model").unwrap(); + let texts = vec!["a".to_string(), "b".to_string()]; + + let result = embed_batch(&client, &store, "test-model", &texts); + handle.join().unwrap(); + + assert_eq!(result.unwrap(), 2, "should return count of stored vectors"); + + let va = store.get("a").expect("vector for 'a' must be cached"); + let vb = store.get("b").expect("vector for 'b' must be cached"); + assert_eq!(va, vec![0.1_f32, 0.2_f32], "vector for 'a' mismatch"); + assert_eq!(vb, vec![0.3_f32, 0.4_f32], "vector for 'b' mismatch"); +} + +#[test] +fn embed_batch_length_mismatch_is_err_no_partial_write() { + let token_dir = tempfile::tempdir().unwrap(); + let embed_dir = tempfile::tempdir().unwrap(); + + // API returns only 1 item for 2 texts + let body = r#"{"data":[{"index":0,"embedding":[0.1,0.2]}]}"#; + let (client, _addr, handle) = make_mock_client(body, &token_dir); + + let store = EmbeddingStore::open(embed_dir.path().to_str().unwrap(), "test-model").unwrap(); + let texts = vec!["a".to_string(), "b".to_string()]; + + let result = embed_batch(&client, &store, "test-model", &texts); + handle.join().unwrap(); + + assert!(result.is_err(), "length mismatch must be an Err, not Ok"); + assert!(store.get("b").is_none(), "partial write for 'b' must not occur"); +} + +#[test] +fn embed_batch_reordered_index_maps_to_correct_text() { + let token_dir = tempfile::tempdir().unwrap(); + let embed_dir = tempfile::tempdir().unwrap(); + + // API returns items in reverse index order (index 1 first, index 0 second) + let body = r#"{"data":[{"index":1,"embedding":[0.3,0.4]},{"index":0,"embedding":[0.1,0.2]}]}"#; + let (client, _addr, handle) = make_mock_client(body, &token_dir); + + let store = EmbeddingStore::open(embed_dir.path().to_str().unwrap(), "test-model").unwrap(); + let texts = vec!["a".to_string(), "b".to_string()]; + + let result = embed_batch(&client, &store, "test-model", &texts); + handle.join().unwrap(); + + assert_eq!(result.unwrap(), 2, "should return count of stored vectors"); + + let va = store.get("a").expect("vector for 'a' must be cached"); + let vb = store.get("b").expect("vector for 'b' must be cached"); + // index 0 → texts[0] = "a" → [0.1, 0.2]; index 1 → texts[1] = "b" → [0.3, 0.4] + assert_eq!(va, vec![0.1_f32, 0.2_f32], "reordered index: vector for 'a' mismatch"); + assert_eq!(vb, vec![0.3_f32, 0.4_f32], "reordered index: vector for 'b' mismatch"); +}