fix: embed_batch length+index pairing; reject non-numeric components
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+16
-2
@@ -72,12 +72,26 @@ pub fn embed_batch(
|
|||||||
let resp: serde_json::Value = client.post_json("/embeddings", &body)?;
|
let resp: serde_json::Value = client.post_json("/embeddings", &body)?;
|
||||||
let data = resp["data"].as_array()
|
let data = resp["data"].as_array()
|
||||||
.ok_or_else(|| AppError::Ionos("embeddings: no data[]".into()))?;
|
.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;
|
let mut n = 0;
|
||||||
for (i, item) in data.iter().enumerate() {
|
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<f32> = item["embedding"].as_array()
|
let v: Vec<f32> = item["embedding"].as_array()
|
||||||
.ok_or_else(|| AppError::Ionos("embeddings: no embedding[]".into()))?
|
.ok_or_else(|| AppError::Ionos("embeddings: no embedding[]".into()))?
|
||||||
.iter().map(|x| x.as_f64().unwrap_or(0.0) as f32).collect();
|
.iter()
|
||||||
store.put(&texts[i], &v)?;
|
.map(|x| x.as_f64().ok_or_else(|| AppError::Ionos("embeddings: non-numeric component".into())))
|
||||||
|
.collect::<Result<Vec<f64>, _>>()?
|
||||||
|
.into_iter().map(|f| f as f32).collect();
|
||||||
|
store.put(&texts[idx], &v)?;
|
||||||
n += 1;
|
n += 1;
|
||||||
}
|
}
|
||||||
Ok(n)
|
Ok(n)
|
||||||
|
|||||||
+102
-1
@@ -1,4 +1,5 @@
|
|||||||
use alpha_id::embed::{EmbeddingStore, estimate};
|
use alpha_id::embed::{EmbeddingStore, embed_batch, estimate};
|
||||||
|
use alpha_id::ionos::IonosClient;
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn cache_key_is_stable_sha256() {
|
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.items, 2);
|
||||||
assert_eq!(est.chars, 5);
|
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::<tiny_http::Header>().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");
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user