use std::sync::Arc;
use serde_json::json;
use tinyagents::harness::embeddings::{
EmbeddingModel, InMemoryVectorStore, MockEmbeddingModel, Retriever, VectorStore,
cosine_similarity,
};
#[test]
fn cosine_similarity_covers_aligned_orthogonal_and_opposite() {
assert_eq!(cosine_similarity(&[1.0, 0.0], &[1.0, 0.0]), 1.0);
assert_eq!(cosine_similarity(&[1.0, 0.0], &[0.0, 1.0]), 0.0);
assert_eq!(cosine_similarity(&[1.0, 0.0], &[-1.0, 0.0]), -1.0);
assert_eq!(cosine_similarity(&[2.0, 0.0], &[5.0, 0.0]), 1.0);
}
#[test]
fn cosine_similarity_degenerate_inputs_return_zero_not_nan() {
assert_eq!(cosine_similarity(&[1.0, 2.0], &[1.0]), 0.0);
assert_eq!(cosine_similarity(&[], &[]), 0.0);
assert_eq!(cosine_similarity(&[0.0, 0.0], &[1.0, 1.0]), 0.0);
}
#[tokio::test]
async fn mock_embedding_is_deterministic_and_correctly_shaped() {
let model = MockEmbeddingModel::new(16);
assert_eq!(model.dimensions(), 16);
let a = model.embed(&["hello".to_string()]).await.unwrap();
let b = model.embed(&["hello".to_string()]).await.unwrap();
assert_eq!(a.len(), 1);
assert_eq!(a[0].len(), 16);
assert_eq!(a, b);
let c = model.embed(&["world".to_string()]).await.unwrap();
assert_ne!(a[0], c[0]);
assert!((cosine_similarity(&a[0], &b[0]) - 1.0).abs() < 1e-6);
}
#[tokio::test]
async fn mock_embedding_batches_preserve_order() {
let model = MockEmbeddingModel::new(8);
let vectors = model
.embed(&["one".to_string(), "two".to_string(), "three".to_string()])
.await
.unwrap();
assert_eq!(vectors.len(), 3);
let two = model.embed(&["two".to_string()]).await.unwrap();
assert_eq!(vectors[1], two[0]);
}
#[tokio::test]
async fn vector_store_add_and_query_by_cosine() {
let store = InMemoryVectorStore::new();
assert!(store.is_empty());
store
.add("a".into(), vec![1.0, 0.0], json!({"t": "x"}))
.await
.unwrap();
store
.add("b".into(), vec![0.0, 1.0], json!({}))
.await
.unwrap();
assert_eq!(store.len(), 2);
let hits = store.query(&[1.0, 0.0], 1).await.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].id, "a");
assert_eq!(hits[0].metadata, json!({"t": "x"}));
assert!((hits[0].score - 1.0).abs() < 1e-6);
}
#[tokio::test]
async fn vector_store_upserts_by_id_in_place() {
let store = InMemoryVectorStore::new();
store
.add("a".into(), vec![1.0, 0.0], json!({"v": 1}))
.await
.unwrap();
store
.add("a".into(), vec![0.0, 1.0], json!({"v": 2}))
.await
.unwrap();
assert_eq!(store.len(), 1);
let hits = store.query(&[0.0, 1.0], 1).await.unwrap();
assert_eq!(hits[0].id, "a");
assert_eq!(hits[0].metadata, json!({"v": 2}));
}
#[tokio::test]
async fn vector_store_rejects_zero_length_and_mismatched_dimensions() {
let store = InMemoryVectorStore::new();
assert!(store.add("z".into(), vec![], json!({})).await.is_err());
store
.add("a".into(), vec![1.0, 0.0], json!({}))
.await
.unwrap();
assert!(
store
.add("b".into(), vec![1.0, 0.0, 0.0], json!({}))
.await
.is_err()
);
assert!(store.query(&[1.0], 1).await.is_err());
}
#[tokio::test]
async fn empty_store_answers_any_query_with_no_hits() {
let store = InMemoryVectorStore::new();
assert!(store.query(&[1.0, 2.0, 3.0], 5).await.unwrap().is_empty());
store.add("a".into(), vec![1.0], json!({})).await.unwrap();
assert!(store.query(&[1.0], 0).await.unwrap().is_empty());
}
#[tokio::test]
async fn retriever_surfaces_dimension_mismatch_from_a_different_model() {
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
let index_retriever = Retriever::new(Arc::new(MockEmbeddingModel::new(32)), store.clone());
index_retriever
.index(vec![("d1".into(), "cats".into(), json!({}))])
.await
.unwrap();
let query_retriever = Retriever::new(Arc::new(MockEmbeddingModel::new(16)), store);
assert!(query_retriever.retrieve("cats", 1).await.is_err());
}