diff --git a/application/prompt_client/prompt_client.py b/application/prompt_client/prompt_client.py index d9e151118..1f27b276e 100644 --- a/application/prompt_client/prompt_client.py +++ b/application/prompt_client/prompt_client.py @@ -1194,13 +1194,14 @@ def get_id_of_most_similar_cre_paginated( most_similar_index = 0 most_similar_id = "" for page in range(starting_page, total_pages + 1): - existing_cres, existing_cre_ids = self.__load_cre_embeddings(embeddings) - - similarities = cosine_similarity(embedding_array, existing_cres) - if np.max(similarities) > max_similarity: - max_similarity = np.max(similarities) - most_similar_index = np.argmax(similarities) - most_similar_id = existing_cre_ids[most_similar_index] + if embeddings: + existing_cres, existing_cre_ids = self.__load_cre_embeddings(embeddings) + + similarities = cosine_similarity(embedding_array, existing_cres) + if np.max(similarities) > max_similarity: + max_similarity = np.max(similarities) + most_similar_index = np.argmax(similarities) + most_similar_id = existing_cre_ids[most_similar_index] if page < total_pages: ( embeddings, @@ -1256,14 +1257,15 @@ def get_id_of_most_similar_node_paginated( most_similar_index = 0 most_similar_id = "" for page in range(starting_page, total_pages + 1): - existing_standards, existing_standard_ids = self.__load_node_embeddings( - embeddings - ) - similarities = cosine_similarity(embedding_array, existing_standards) - if np.max(similarities) > max_similarity: - max_similarity = np.max(similarities) - most_similar_index = int(np.argmax(similarities)) - most_similar_id = existing_standard_ids[most_similar_index] + if embeddings: + existing_standards, existing_standard_ids = self.__load_node_embeddings( + embeddings + ) + similarities = cosine_similarity(embedding_array, existing_standards) + if np.max(similarities) > max_similarity: + max_similarity = np.max(similarities) + most_similar_index = int(np.argmax(similarities)) + most_similar_id = existing_standard_ids[most_similar_index] if page < total_pages: embeddings, _, _ = self.database.get_embeddings_by_doc_type_paginated( diff --git a/application/tests/prompt_client_pgvector_similarity_test.py b/application/tests/prompt_client_pgvector_similarity_test.py index d6e6a32fe..af4f3bee0 100644 --- a/application/tests/prompt_client_pgvector_similarity_test.py +++ b/application/tests/prompt_client_pgvector_similarity_test.py @@ -149,6 +149,36 @@ def test_cre_paginated_single_page(self) -> None: self.assertEqual(result, ("only-cre", 1.0)) self.assertEqual(database.get_embeddings_by_doc_type_paginated.call_count, 1) + def test_node_paginated_skips_empty_final_page(self) -> None: + # The final page's embeddings all failed to parse and came back + # empty (e.g. malformed stored vectors). Must not crash + # cosine_similarity with a feature-count mismatch; the real match + # found on an earlier page must still be returned. + pages = { + 1: {"target-node": self.MATCHING_VECTOR}, + 2: {}, + } + handler, database = self._make_handler(can_use_pgvector=False, pages=pages) + + result = handler.get_id_of_most_similar_node_paginated( + self.QUERY_EMBEDDING, similarity_threshold=0.5 + ) + + self.assertEqual(result, ("target-node", 1.0)) + + def test_cre_paginated_skips_empty_final_page(self) -> None: + pages = { + 1: {"target-cre": self.MATCHING_VECTOR}, + 2: {}, + } + handler, database = self._make_handler(can_use_pgvector=False, pages=pages) + + result = handler.get_id_of_most_similar_cre_paginated( + self.QUERY_EMBEDDING, similarity_threshold=0.5 + ) + + self.assertEqual(result, ("target-cre", 1.0)) + class FindMostSimilarEmbeddingIdResilienceTest(unittest.TestCase): def test_query_error_returns_no_match(self) -> None: