""" FAISS向量管理器测试文件 用于验证FAISSManager的各项功能 """ import unittest import numpy as np import tempfile import shutil import os import sys # 添加项目根目录到路径 sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from faiss_vector_db.faiss_manager import FAISSManager class TestFAISSManager(unittest.TestCase): """FAISS管理器测试类""" def setUp(self): """测试前的设置""" self.dimension = 64 self.manager = FAISSManager( dimension=self.dimension, index_type="IndexFlatIP", normalize_vectors=True, metric_type="cosine" ) # 创建临时目录用于测试保存/加载 self.temp_dir = tempfile.mkdtemp() def tearDown(self): """测试后的清理""" # 删除临时目录 if os.path.exists(self.temp_dir): shutil.rmtree(self.temp_dir) def test_initialization(self): """测试初始化""" self.assertEqual(self.manager.dimension, self.dimension) self.assertEqual(self.manager.index_type, "IndexFlatIP") self.assertTrue(self.manager.normalize_vectors) self.assertEqual(self.manager.metric_type, "cosine") self.assertEqual(self.manager.index.ntotal, 0) def test_add_single_vector(self): """测试添加单个向量,用numpy随机生成一个向量""" vector = np.random.random(self.dimension).astype(np.float32) vector_id = "test_001" metadata = {"name": "测试菜品", "category": "测试"} # 测试添加 result = self.manager.add_vector(vector, vector_id, metadata) self.assertTrue(result) # 验证索引状态 self.assertEqual(self.manager.index.ntotal, 1) self.assertIn(vector_id, self.manager.reverse_id_mapping) self.assertIn(vector_id, self.manager.metadata) # 测试重复添加 result = self.manager.add_vector(vector, vector_id, metadata) self.assertTrue(result) # 应该更新成功 self.assertEqual(self.manager.index.ntotal, 2) # 因为,没有真正的删除 def test_add_batch_vectors(self): """测试批量添加向量""" batch_size = 10 vectors = np.random.random((batch_size, self.dimension)).astype(np.float32) vector_ids = [f"batch_{i}" for i in range(batch_size)] metadata_list = [{"name": f"菜品{i}", "category": "测试"} for i in range(batch_size)] # 测试批量添加 results = self.manager.add_vectors_batch(vectors, vector_ids, metadata_list) # 验证结果 self.assertEqual(len(results), batch_size) self.assertTrue(all(results)) self.assertEqual(self.manager.index.ntotal, batch_size) # 验证所有向量都被添加 for vector_id in vector_ids: self.assertIn(vector_id, self.manager.reverse_id_mapping) self.assertIn(vector_id, self.manager.metadata) def test_search_similar(self): """测试相似度搜索""" # 先添加一些向量 batch_size = 20 vectors = np.random.random((batch_size, self.dimension)).astype(np.float32) vector_ids = [f"search_test_{i}" for i in range(batch_size)] metadata_list = [{"name": f"菜品{i}", "score": i} for i in range(batch_size)] self.manager.add_vectors_batch(vectors, vector_ids, metadata_list) # 测试搜索 query_vector = np.random.random(self.dimension).astype(np.float32) k = 5 results = self.manager.search_similar(query_vector, k=k) # 验证结果 self.assertLessEqual(len(results), k) self.assertLessEqual(len(results), batch_size) # 验证结果格式 for vector_id, similarity, metadata in results: self.assertIsInstance(vector_id, str) self.assertIsInstance(similarity, float) self.assertIsInstance(metadata, dict) self.assertIn(vector_id, self.manager.reverse_id_mapping) # 测试带阈值的搜索 threshold = 0.8 results_with_threshold = self.manager.search_similar( query_vector, k=k, threshold=threshold ) # 验证阈值过滤 for _, similarity, _ in results_with_threshold: self.assertGreaterEqual(similarity, threshold) def test_get_vector_by_id(self): """测试根据ID获取向量""" vector = np.random.random(self.dimension).astype(np.float32) vector_id = "get_test_001" metadata = {"name": "测试获取", "value": 42} # 添加向量 self.manager.add_vector(vector, vector_id, metadata) # 测试获取存在的向量 result = self.manager.get_vector_by_id(vector_id) self.assertIsNotNone(result) vector_data, retrieved_metadata = result self.assertEqual(retrieved_metadata, metadata) # 测试获取不存在的向量 result = self.manager.get_vector_by_id("nonexistent") self.assertIsNone(result) def test_update_vector(self): """测试更新向量""" # 添加初始向量 vector = np.random.random(self.dimension).astype(np.float32) vector_id = "update_test_001" metadata = {"name": "原始菜品", "version": 1} self.manager.add_vector(vector, vector_id, metadata) initial_count = self.manager.index.ntotal # 更新向量 new_vector = np.random.random(self.dimension).astype(np.float32) new_metadata = {"name": "更新菜品", "version": 2} result = self.manager.update_vector(new_vector, vector_id, new_metadata) self.assertTrue(result) # 验证更新后的状态 self.assertEqual(self.manager.index.ntotal, initial_count) # 验证元数据已更新 vector_info = self.manager.get_vector_by_id(vector_id) self.assertIsNotNone(vector_info) _, retrieved_metadata = vector_info self.assertEqual(retrieved_metadata, new_metadata) # 测试更新不存在的向量 result = self.manager.update_vector(new_vector, "nonexistent", new_metadata) self.assertFalse(result) def test_delete_vector(self): """测试删除向量""" # 添加向量 vector = np.random.random(self.dimension).astype(np.float32) vector_id = "delete_test_001" metadata = {"name": "待删除菜品"} self.manager.add_vector(vector, vector_id, metadata) # 验证向量存在 self.assertIn(vector_id, self.manager.reverse_id_mapping) self.assertIn(vector_id, self.manager.metadata) # 删除向量 result = self.manager.delete_vector(vector_id) self.assertTrue(result) # 验证向量已删除 self.assertNotIn(vector_id, self.manager.reverse_id_mapping) self.assertNotIn(vector_id, self.manager.metadata) # 测试删除不存在的向量 result = self.manager.delete_vector("nonexistent") self.assertFalse(result) def test_delete_vectors_batch(self): """测试批量删除向量""" # 添加多个向量 batch_size = 5 vectors = np.random.random((batch_size, self.dimension)).astype(np.float32) vector_ids = [f"batch_delete_{i}" for i in range(batch_size)] self.manager.add_vectors_batch(vectors, vector_ids) # 批量删除部分向量 delete_ids = vector_ids[:3] results = self.manager.delete_vectors_batch(delete_ids) # 验证删除结果 self.assertEqual(len(results), len(delete_ids)) self.assertTrue(all(results)) # 验证向量已删除 for vector_id in delete_ids: self.assertNotIn(vector_id, self.manager.reverse_id_mapping) # 验证剩余向量仍存在 for vector_id in vector_ids[3:]: self.assertIn(vector_id, self.manager.reverse_id_mapping) def test_save_and_load_index(self): """测试保存和加载索引""" # 添加一些向量 batch_size = 10 vectors = np.random.random((batch_size, self.dimension)).astype(np.float32) vector_ids = [f"save_load_{i}" for i in range(batch_size)] metadata_list = [{"name": f"菜品{i}", "id": i} for i in range(batch_size)] self.manager.add_vectors_batch(vectors, vector_ids, metadata_list) # 保存索引 save_path = os.path.join(self.temp_dir, "test_index") result = self.manager.save_index(save_path) self.assertTrue(result) # 验证文件已创建 self.assertTrue(os.path.exists(os.path.join(save_path, "faiss_index.bin"))) self.assertTrue(os.path.exists(os.path.join(save_path, "metadata.pkl"))) self.assertTrue(os.path.exists(os.path.join(save_path, "config.json"))) # 创建新管理器并加载索引 new_manager = FAISSManager(dimension=self.dimension) result = new_manager.load_index(save_path) self.assertTrue(result) # 验证加载后的状态 self.assertEqual(new_manager.index.ntotal, batch_size) self.assertEqual(len(new_manager.reverse_id_mapping), batch_size) self.assertEqual(len(new_manager.metadata), batch_size) # 验证数据一致性 for vector_id in vector_ids: self.assertIn(vector_id, new_manager.reverse_id_mapping) self.assertIn(vector_id, new_manager.metadata) # 测试加载后的搜索功能 query_vector = np.random.random(self.dimension).astype(np.float32) results = new_manager.search_similar(query_vector, k=3) self.assertGreater(len(results), 0) def test_clear_index(self): """测试清空索引""" # 添加一些向量 vectors = np.random.random((5, self.dimension)).astype(np.float32) vector_ids = [f"clear_test_{i}" for i in range(5)] self.manager.add_vectors_batch(vectors, vector_ids) # 验证向量已添加 self.assertEqual(self.manager.index.ntotal, 5) self.assertEqual(len(self.manager.reverse_id_mapping), 5) # 清空索引 result = self.manager.clear_index() self.assertTrue(result) # 验证索引已清空 self.assertEqual(self.manager.index.ntotal, 0) self.assertEqual(len(self.manager.reverse_id_mapping), 0) self.assertEqual(len(self.manager.metadata), 0) self.assertEqual(self.manager.next_faiss_id, 0) def test_get_index_info(self): """测试获取索引信息""" # 初始状态 info = self.manager.get_index_info() self.assertEqual(info['total_vectors'], 0) self.assertEqual(info['dimension'], self.dimension) self.assertEqual(info['active_vectors'], 0) # 添加向量后 vectors = np.random.random((3, self.dimension)).astype(np.float32) vector_ids = [f"info_test_{i}" for i in range(3)] self.manager.add_vectors_batch(vectors, vector_ids) info = self.manager.get_index_info() self.assertEqual(info['total_vectors'], 3) self.assertEqual(info['active_vectors'], 3) # 删除一个向量后 self.manager.delete_vector(vector_ids[0]) info = self.manager.get_index_info() self.assertEqual(info['total_vectors'], 3) # FAISS索引中的总数不变 self.assertEqual(info['active_vectors'], 2) # 活跃向量数减少 def test_different_index_types(self): """测试不同的索引类型""" index_types = ["IndexFlatIP", "IndexFlatL2"] for index_type in index_types: with self.subTest(index_type=index_type): manager = FAISSManager( dimension=self.dimension, index_type=index_type ) # 添加向量 vector = np.random.random(self.dimension).astype(np.float32) result = manager.add_vector(vector, f"test_{index_type}") self.assertTrue(result) # 搜索 query = np.random.random(self.dimension).astype(np.float32) results = manager.search_similar(query, k=1) self.assertEqual(len(results), 1) def test_vector_normalization(self): """测试向量标准化""" # 测试标准化开启 manager_norm = FAISSManager( dimension=self.dimension, normalize_vectors=True ) # 测试标准化关闭 manager_no_norm = FAISSManager( dimension=self.dimension, normalize_vectors=False ) # 添加相同的向量 vector = np.random.random(self.dimension).astype(np.float32) * 10 # 放大向量 manager_norm.add_vector(vector, "norm_test") manager_no_norm.add_vector(vector, "no_norm_test") # 两个管理器都应该成功添加 self.assertEqual(manager_norm.index.ntotal, 1) self.assertEqual(manager_no_norm.index.ntotal, 1) class TestFAISSManagerIntegration(unittest.TestCase): """FAISS管理器集成测试""" def test_food_classification_scenario(self): """测试食物分类场景""" # 创建管理器 manager = FAISSManager(dimension=128, index_type="IndexFlatIP") # 模拟食物向量数据 food_data = [ {"id": "sichuan_001", "name": "回锅肉", "category": "川菜", "spicy": 3}, {"id": "sichuan_002", "name": "麻辣小面", "category": "川菜", "spicy": 4}, {"id": "sichuan_003", "name": "宫保鸡丁", "category": "川菜", "spicy": 3}, {"id": "home_001", "name": "西红柿鸡蛋", "category": "家常菜", "spicy": 0}, {"id": "home_002", "name": "炒细面", "category": "家常菜", "spicy": 1}, ] # 生成模拟向量(实际应用中这些是通过embedding模型生成的) vectors = [] for i, food in enumerate(food_data): # 为川菜生成相似的向量,为家常菜生成另一类相似的向量 if food["category"] == "川菜": base_vector = np.array([1.0] * 64 + [0.0] * 64) else: base_vector = np.array([0.0] * 64 + [1.0] * 64) # 添加随机噪声 noise = np.random.normal(0, 0.1, 128) vector = (base_vector + noise).astype(np.float32) vectors.append(vector) # 批量添加向量 vector_ids = [food["id"] for food in food_data] metadata_list = [ {"name": food["name"], "category": food["category"], "spicy": food["spicy"]} for food in food_data ] results = manager.add_vectors_batch(vectors, vector_ids, metadata_list) self.assertTrue(all(results)) # 测试川菜查询 sichuan_query = np.array([1.0] * 64 + [0.0] * 64).astype(np.float32) sichuan_results = manager.search_similar(sichuan_query, k=3) # 验证川菜结果 self.assertGreater(len(sichuan_results), 0) for vector_id, similarity, metadata in sichuan_results: self.assertEqual(metadata["category"], "川菜") # 测试家常菜查询 home_query = np.array([0.0] * 64 + [1.0] * 64).astype(np.float32) home_results = manager.search_similar(home_query, k=3) # 验证家常菜结果 self.assertGreater(len(home_results), 0) for vector_id, similarity, metadata in home_results: self.assertEqual(metadata["category"], "家常菜") # 测试保存和加载 temp_dir = tempfile.mkdtemp() try: save_path = os.path.join(temp_dir, "food_index") self.assertTrue(manager.save_index(save_path)) # 加载到新管理器 new_manager = FAISSManager(dimension=128) self.assertTrue(new_manager.load_index(save_path)) # 验证加载后的搜索功能 loaded_results = new_manager.search_similar(sichuan_query, k=2) self.assertGreater(len(loaded_results), 0) finally: shutil.rmtree(temp_dir) def run_tests(): """运行所有测试""" # 创建测试套件 test_suite = unittest.TestSuite() # 添加基本功能测试 test_suite.addTest(unittest.makeSuite(TestFAISSManager)) # 添加集成测试 test_suite.addTest(unittest.makeSuite(TestFAISSManagerIntegration)) # 运行测试 runner = unittest.TextTestRunner(verbosity=2) result = runner.run(test_suite) return result.wasSuccessful() if __name__ == "__main__": print("开始运行FAISS管理器测试...") success = run_tests() if success: print("\n✅ 所有测试通过!") else: print("\n❌ 部分测试失败,请检查代码。") sys.exit(1)