feat: add a unified local vector store (#6995)
* add a unified local vector store * [autofix.ci] apply automated fixes * fixed lint Error * [autofix.ci] apply automated fixes * Update src/backend/base/langflow/components/vectorstores/local_db.py Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> * refactor: Enhance type hints and clean up imports in LocalDBComponent - Added type hints for the `update_build_config` method parameters and return type. - Removed unused import of `override`. - Cleaned up the `build_vector_store` method by removing the import error handling for `Chroma`, as it is now assumed to be handled elsewhere. * test: Add unit tests for LocalDBComponent functionality - Introduced comprehensive tests for the LocalDBComponent, covering database creation, data ingestion, similarity search, and duplicate handling. - Implemented fixtures for default parameters and collection mappings. - Verified the behavior of various search types and ensured correct handling of duplicates. - Added tests for configuration updates and listing existing collections. * feat: Implement equality comparison for DataFrame class - Added an __eq__ method to the DataFrame class to handle comparisons with empty DataFrames and non-DataFrame objects. - Ensures that empty DataFrames and empty lists are treated as unequal, improving the robustness of DataFrame comparisons. * Update local_db.py * [autofix.ci] apply automated fixes * removed data output * [autofix.ci] apply automated fixes --------- Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Edwin Jose <edwin.jose@datastax.com> Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> Co-authored-by: Gabriel Luiz Freitas Almeida <gabriel@langflow.org> Co-authored-by: Ítalo Johnny <italojohnnydosanjos@gmail.com>
This commit is contained in:
parent
57524d538c
commit
fecd2b5598
3 changed files with 639 additions and 0 deletions
|
|
@ -9,6 +9,7 @@ from .elasticsearch import ElasticsearchVectorStoreComponent
|
|||
from .faiss import FaissVectorStoreComponent
|
||||
from .graph_rag import GraphRAGComponent
|
||||
from .hcd import HCDVectorStoreComponent
|
||||
from .local_db import LocalDBComponent
|
||||
from .milvus import MilvusVectorStoreComponent
|
||||
from .mongodb_atlas import MongoVectorStoreComponent
|
||||
from .opensearch import OpenSearchVectorStoreComponent
|
||||
|
|
@ -35,6 +36,7 @@ __all__ = [
|
|||
"FaissVectorStoreComponent",
|
||||
"GraphRAGComponent",
|
||||
"HCDVectorStoreComponent",
|
||||
"LocalDBComponent",
|
||||
"MilvusVectorStoreComponent",
|
||||
"MongoVectorStoreComponent",
|
||||
"OpenSearchVectorStoreComponent",
|
||||
|
|
|
|||
255
src/backend/base/langflow/components/vectorstores/local_db.py
Normal file
255
src/backend/base/langflow/components/vectorstores/local_db.py
Normal file
|
|
@ -0,0 +1,255 @@
|
|||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
|
||||
from langchain_chroma import Chroma
|
||||
from loguru import logger
|
||||
from typing_extensions import override
|
||||
|
||||
from langflow.base.vectorstores.model import LCVectorStoreComponent, check_cached_vector_store
|
||||
from langflow.base.vectorstores.utils import chroma_collection_to_data
|
||||
from langflow.inputs.inputs import MultilineInput
|
||||
from langflow.io import BoolInput, DropdownInput, HandleInput, IntInput, MessageTextInput, TabInput
|
||||
from langflow.schema import Data, DataFrame
|
||||
from langflow.template.field.base import Output
|
||||
|
||||
|
||||
class LocalDBComponent(LCVectorStoreComponent):
|
||||
"""Chroma Vector Store with search capabilities."""
|
||||
|
||||
display_name: str = "Local DB"
|
||||
description: str = "Local Vector Store with search capabilities"
|
||||
name = "LocalDB"
|
||||
icon = "database"
|
||||
|
||||
inputs = [
|
||||
TabInput(
|
||||
name="mode",
|
||||
display_name="Mode",
|
||||
options=["Ingest", "Retrieve"],
|
||||
info="Select the operation mode",
|
||||
value="Ingest",
|
||||
real_time_refresh=True,
|
||||
show=True,
|
||||
),
|
||||
MessageTextInput(
|
||||
name="collection_name",
|
||||
display_name="Collection Name",
|
||||
value="langflow",
|
||||
),
|
||||
MessageTextInput(
|
||||
name="persist_directory",
|
||||
display_name="Persist Directory",
|
||||
info=(
|
||||
"Custom base directory to save the vector store. "
|
||||
"Collections will be stored under '{directory}/vector_stores/{collection_name}'. "
|
||||
"If not specified, it will use your system's cache folder."
|
||||
),
|
||||
advanced=True,
|
||||
),
|
||||
DropdownInput(
|
||||
name="existing_collections",
|
||||
display_name="Existing Collections",
|
||||
options=[], # Will be populated dynamically
|
||||
info="Select a previously created collection to search through its stored data.",
|
||||
show=False,
|
||||
combobox=True,
|
||||
),
|
||||
HandleInput(name="embedding", display_name="Embedding", input_types=["Embeddings"]),
|
||||
BoolInput(
|
||||
name="allow_duplicates",
|
||||
display_name="Allow Duplicates",
|
||||
advanced=True,
|
||||
info="If false, will not add documents that are already in the Vector Store.",
|
||||
),
|
||||
DropdownInput(
|
||||
name="search_type",
|
||||
display_name="Search Type",
|
||||
options=["Similarity", "MMR"],
|
||||
value="Similarity",
|
||||
advanced=True,
|
||||
),
|
||||
HandleInput(
|
||||
name="ingest_data",
|
||||
display_name="Ingest Data",
|
||||
input_types=["Data", "DataFrame"],
|
||||
is_list=True,
|
||||
info="Data to store. It will be embedded and indexed for semantic search.",
|
||||
show=True,
|
||||
),
|
||||
MultilineInput(
|
||||
name="search_query",
|
||||
display_name="Search Query",
|
||||
tool_mode=True,
|
||||
info="Enter text to search for similar content in the selected collection.",
|
||||
show=False,
|
||||
),
|
||||
IntInput(
|
||||
name="number_of_results",
|
||||
display_name="Number of Results",
|
||||
info="Number of results to return.",
|
||||
advanced=True,
|
||||
value=10,
|
||||
),
|
||||
IntInput(
|
||||
name="limit",
|
||||
display_name="Limit",
|
||||
advanced=True,
|
||||
info="Limit the number of records to compare when Allow Duplicates is False.",
|
||||
),
|
||||
]
|
||||
outputs = [
|
||||
Output(display_name="DataFrame", name="dataframe", method="as_dataframe"),
|
||||
]
|
||||
|
||||
def get_vector_store_directory(self, base_dir: str | Path) -> Path:
|
||||
"""Get the full directory path for a collection."""
|
||||
# Ensure base_dir is a Path object
|
||||
base_dir = Path(base_dir)
|
||||
# Create the full path: base_dir/vector_stores/collection_name
|
||||
full_path = base_dir / "vector_stores" / self.collection_name
|
||||
# Create the directory if it doesn't exist
|
||||
full_path.mkdir(parents=True, exist_ok=True)
|
||||
return full_path
|
||||
|
||||
def get_default_persist_dir(self) -> str:
|
||||
"""Get the default persist directory from cache."""
|
||||
from langflow.services.cache.utils import CACHE_DIR
|
||||
|
||||
return str(self.get_vector_store_directory(CACHE_DIR))
|
||||
|
||||
def list_existing_collections(self) -> list[str]:
|
||||
"""List existing vector store collections from the persist directory."""
|
||||
from langflow.services.cache.utils import CACHE_DIR
|
||||
|
||||
# Get the base directory (either custom or cache)
|
||||
base_dir = Path(self.persist_directory) if self.persist_directory else Path(CACHE_DIR)
|
||||
# Get the vector_stores subdirectory
|
||||
vector_stores_dir = base_dir / "vector_stores"
|
||||
if not vector_stores_dir.exists():
|
||||
return []
|
||||
|
||||
return [d.name for d in vector_stores_dir.iterdir() if d.is_dir()]
|
||||
|
||||
def update_build_config(self, build_config: dict, field_value: str, field_name: str | None = None) -> dict:
|
||||
"""Update the build configuration when the mode changes."""
|
||||
if field_name == "mode":
|
||||
# Hide all dynamic fields by default
|
||||
dynamic_fields = [
|
||||
"ingest_data",
|
||||
"search_query",
|
||||
"search_type",
|
||||
"number_of_results",
|
||||
"existing_collections",
|
||||
"collection_name",
|
||||
"embedding",
|
||||
"allow_duplicates",
|
||||
"limit",
|
||||
]
|
||||
for field in dynamic_fields:
|
||||
if field in build_config:
|
||||
build_config[field]["show"] = False
|
||||
|
||||
# Show/hide fields based on selected mode
|
||||
if field_value == "Ingest":
|
||||
if "ingest_data" in build_config:
|
||||
build_config["ingest_data"]["show"] = True
|
||||
if "collection_name" in build_config:
|
||||
build_config["collection_name"]["show"] = True
|
||||
build_config["collection_name"]["display_name"] = "Name Your Collection"
|
||||
if "persist" in build_config:
|
||||
build_config["persist"]["show"] = True
|
||||
if "persist_directory" in build_config:
|
||||
build_config["persist_directory"]["show"] = True
|
||||
if "embedding" in build_config:
|
||||
build_config["embedding"]["show"] = True
|
||||
if "allow_duplicates" in build_config:
|
||||
build_config["allow_duplicates"]["show"] = True
|
||||
if "limit" in build_config:
|
||||
build_config["limit"]["show"] = True
|
||||
elif field_value == "Retrieve":
|
||||
if "persist" in build_config:
|
||||
build_config["persist"]["show"] = False
|
||||
build_config["search_query"]["show"] = True
|
||||
build_config["search_type"]["show"] = True
|
||||
build_config["number_of_results"]["show"] = True
|
||||
build_config["embedding"]["show"] = True
|
||||
build_config["collection_name"]["show"] = False
|
||||
# Show existing collections dropdown and update its options
|
||||
if "existing_collections" in build_config:
|
||||
build_config["existing_collections"]["show"] = True
|
||||
build_config["existing_collections"]["options"] = self.list_existing_collections()
|
||||
# Hide collection_name in Retrieve mode since we use existing_collections
|
||||
elif field_name == "existing_collections":
|
||||
# Update collection_name when an existing collection is selected
|
||||
if "collection_name" in build_config:
|
||||
build_config["collection_name"]["value"] = field_value
|
||||
|
||||
return build_config
|
||||
|
||||
@override
|
||||
@check_cached_vector_store
|
||||
def build_vector_store(self) -> Chroma:
|
||||
"""Builds the Chroma object."""
|
||||
try:
|
||||
from langchain_chroma import Chroma
|
||||
except ImportError as e:
|
||||
msg = "Could not import Chroma integration package. Please install it with `pip install langchain-chroma`."
|
||||
raise ImportError(msg) from e
|
||||
# Chroma settings
|
||||
# chroma_settings = None
|
||||
if self.existing_collections:
|
||||
self.collection_name = self.existing_collections
|
||||
|
||||
# Use user-provided directory or default cache directory
|
||||
if self.persist_directory:
|
||||
base_dir = self.resolve_path(self.persist_directory)
|
||||
persist_directory = str(self.get_vector_store_directory(base_dir))
|
||||
logger.debug(f"Using custom persist directory: {persist_directory}")
|
||||
else:
|
||||
persist_directory = self.get_default_persist_dir()
|
||||
logger.debug(f"Using default persist directory: {persist_directory}")
|
||||
|
||||
chroma = Chroma(
|
||||
persist_directory=persist_directory,
|
||||
client=None,
|
||||
embedding_function=self.embedding,
|
||||
collection_name=self.collection_name,
|
||||
)
|
||||
|
||||
self._add_documents_to_vector_store(chroma)
|
||||
self.status = chroma_collection_to_data(chroma.get(limit=self.limit))
|
||||
return chroma
|
||||
|
||||
def _add_documents_to_vector_store(self, vector_store: "Chroma") -> None:
|
||||
"""Adds documents to the Vector Store."""
|
||||
ingest_data: list | Data | DataFrame = self.ingest_data
|
||||
if not ingest_data:
|
||||
self.status = ""
|
||||
return
|
||||
|
||||
# Convert DataFrame to Data if needed using parent's method
|
||||
ingest_data = self._prepare_ingest_data()
|
||||
|
||||
stored_documents_without_id = []
|
||||
if self.allow_duplicates:
|
||||
stored_data = []
|
||||
else:
|
||||
stored_data = chroma_collection_to_data(vector_store.get(limit=self.limit))
|
||||
for value in deepcopy(stored_data):
|
||||
del value.id
|
||||
stored_documents_without_id.append(value)
|
||||
|
||||
documents = []
|
||||
for _input in ingest_data or []:
|
||||
if isinstance(_input, Data):
|
||||
if _input not in stored_documents_without_id:
|
||||
documents.append(_input.to_lc_document())
|
||||
else:
|
||||
msg = "Vector Store Inputs must be Data objects."
|
||||
raise TypeError(msg)
|
||||
|
||||
if documents and self.embedding is not None:
|
||||
self.log(f"Adding {len(documents)} documents to the Vector Store.")
|
||||
vector_store.add_documents(documents)
|
||||
else:
|
||||
self.log("No documents to add to the Vector Store.")
|
||||
|
|
@ -0,0 +1,382 @@
|
|||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from langflow.components.vectorstores.local_db import LocalDBComponent
|
||||
from langflow.schema.data import Data
|
||||
from langflow.services.cache.utils import CACHE_DIR
|
||||
|
||||
from tests.base import ComponentTestBaseWithoutClient, VersionComponentMapping
|
||||
|
||||
|
||||
@pytest.mark.api_key_required
|
||||
class TestLocalDBComponent(ComponentTestBaseWithoutClient):
|
||||
@pytest.fixture
|
||||
def component_class(self) -> type[Any]:
|
||||
"""Return the component class to test."""
|
||||
return LocalDBComponent
|
||||
|
||||
@pytest.fixture
|
||||
def default_kwargs(self, tmp_path: Path) -> dict[str, Any]:
|
||||
"""Return the default kwargs for the component."""
|
||||
from langflow.components.embeddings.openai import OpenAIEmbeddingsComponent
|
||||
|
||||
if os.getenv("OPENAI_API_KEY") is None:
|
||||
pytest.skip("OPENAI_API_KEY is not set")
|
||||
|
||||
api_key = os.getenv("OPENAI_API_KEY")
|
||||
|
||||
return {
|
||||
"embedding": OpenAIEmbeddingsComponent(openai_api_key=api_key).build_embeddings(),
|
||||
"collection_name": "test_collection",
|
||||
"persist": True,
|
||||
"persist_directory": str(tmp_path), # Convert Path to string
|
||||
"mode": "Ingest",
|
||||
}
|
||||
|
||||
@pytest.fixture
|
||||
def file_names_mapping(self) -> list[VersionComponentMapping]:
|
||||
"""Return the file names mapping for different versions."""
|
||||
# Return an empty list since this is a new component
|
||||
return []
|
||||
|
||||
def test_create_db(self, component_class: type[LocalDBComponent], default_kwargs: dict[str, Any]) -> None:
|
||||
"""Test creating a vector store."""
|
||||
component: LocalDBComponent = component_class().set(**default_kwargs)
|
||||
component.build_vector_store()
|
||||
persist_directory = Path(default_kwargs["persist_directory"])
|
||||
assert persist_directory.exists()
|
||||
assert persist_directory.is_dir()
|
||||
# Assert it isn't empty
|
||||
assert len(list(persist_directory.iterdir())) > 0
|
||||
# Assert there's a chroma.sqlite3 file (since LocalDB uses Chroma underneath)
|
||||
assert (persist_directory / "chroma.sqlite3").exists()
|
||||
assert (persist_directory / "chroma.sqlite3").is_file()
|
||||
|
||||
@patch("langchain_chroma.Chroma._collection")
|
||||
def test_create_db_with_data(
|
||||
self,
|
||||
mock_collection,
|
||||
component_class: type[LocalDBComponent],
|
||||
default_kwargs: dict[str, Any],
|
||||
) -> None:
|
||||
"""Test creating a vector store with data."""
|
||||
# Set ingest_data in default_kwargs to a list of Data objects
|
||||
test_texts = ["test data 1", "test data 2", "something completely different"]
|
||||
default_kwargs["ingest_data"] = [Data(text=text) for text in test_texts]
|
||||
|
||||
# Mock the collection count to return the expected number
|
||||
mock_collection.count.return_value = len(test_texts)
|
||||
mock_collection.name = default_kwargs["collection_name"]
|
||||
|
||||
# Mock the _add_documents_to_vector_store method to ensure add_documents is called
|
||||
with patch.object(LocalDBComponent, "_add_documents_to_vector_store") as mock_add_docs_method:
|
||||
component: LocalDBComponent = component_class().set(**default_kwargs)
|
||||
vector_store = component.build_vector_store()
|
||||
|
||||
# Verify the method was called
|
||||
mock_add_docs_method.assert_called_once()
|
||||
|
||||
# Verify collection exists and has the correct data
|
||||
assert vector_store._collection.name == default_kwargs["collection_name"]
|
||||
assert vector_store._collection.count() == len(test_texts)
|
||||
|
||||
def test_default_persist_dir(self, component_class: type[LocalDBComponent], default_kwargs: dict[str, Any]) -> None:
|
||||
"""Test the default persist directory functionality."""
|
||||
# Remove persist_directory from default_kwargs to test default directory
|
||||
default_kwargs.pop("persist_directory")
|
||||
|
||||
component: LocalDBComponent = component_class().set(**default_kwargs)
|
||||
|
||||
# Call get_default_persist_dir and check the result
|
||||
default_dir = component.get_default_persist_dir()
|
||||
expected_dir = Path(CACHE_DIR) / "vector_stores" / default_kwargs["collection_name"]
|
||||
|
||||
assert Path(default_dir) == expected_dir
|
||||
assert Path(default_dir).exists()
|
||||
|
||||
@patch("langchain_chroma.Chroma.similarity_search")
|
||||
def test_similarity_search(
|
||||
self,
|
||||
mock_similarity_search,
|
||||
component_class: type[LocalDBComponent],
|
||||
default_kwargs: dict[str, Any],
|
||||
) -> None:
|
||||
"""Test the similarity search functionality."""
|
||||
# Create test data with distinct topics
|
||||
test_data = [
|
||||
"The quick brown fox jumps over the lazy dog",
|
||||
"Python is a popular programming language",
|
||||
"Machine learning models process data",
|
||||
"The lazy dog sleeps all day long",
|
||||
]
|
||||
default_kwargs["ingest_data"] = [Data(text=text) for text in test_data]
|
||||
default_kwargs["search_type"] = "Similarity"
|
||||
default_kwargs["number_of_results"] = 2
|
||||
|
||||
# Mock the similarity_search to return documents
|
||||
from langchain_core.documents import Document
|
||||
|
||||
mock_docs = [
|
||||
Document(page_content="The lazy dog sleeps all day long"),
|
||||
Document(page_content="The quick brown fox jumps over the lazy dog"),
|
||||
]
|
||||
mock_similarity_search.return_value = mock_docs
|
||||
|
||||
component: LocalDBComponent = component_class().set(**default_kwargs)
|
||||
component.build_vector_store()
|
||||
|
||||
# Switch to Retrieve mode
|
||||
component.set(mode="Retrieve", search_query="dog sleeping")
|
||||
results = component.search_documents()
|
||||
|
||||
assert len(results) == 2
|
||||
# The most relevant results should be about dogs
|
||||
assert any("dog" in result.text.lower() for result in results)
|
||||
mock_similarity_search.assert_called_once_with(query="dog sleeping", k=2)
|
||||
|
||||
# Test with different number of results
|
||||
component.set(number_of_results=3)
|
||||
another_doc = Document(page_content="Another document")
|
||||
mock_similarity_search.return_value = [*mock_docs, another_doc] # Use unpacking instead of concatenation
|
||||
results = component.search_documents()
|
||||
assert len(results) == 3
|
||||
|
||||
@patch("langchain_chroma.Chroma.max_marginal_relevance_search")
|
||||
def test_mmr_search(
|
||||
self,
|
||||
mock_mmr_search,
|
||||
component_class: type[LocalDBComponent],
|
||||
default_kwargs: dict[str, Any],
|
||||
) -> None:
|
||||
"""Test the MMR search functionality."""
|
||||
# Create test data with some similar documents
|
||||
test_data = [
|
||||
"The quick brown fox jumps",
|
||||
"The quick brown fox leaps",
|
||||
"The quick brown fox hops",
|
||||
"Something completely different about cats",
|
||||
]
|
||||
default_kwargs["ingest_data"] = [Data(text=text) for text in test_data]
|
||||
default_kwargs["search_type"] = "MMR"
|
||||
default_kwargs["number_of_results"] = 3
|
||||
|
||||
# Mock the MMR search to return documents
|
||||
from langchain_core.documents import Document
|
||||
|
||||
mock_docs = [
|
||||
Document(page_content="The quick brown fox jumps"),
|
||||
Document(page_content="The quick brown fox leaps"),
|
||||
Document(page_content="Something completely different about cats"),
|
||||
]
|
||||
mock_mmr_search.return_value = mock_docs
|
||||
|
||||
component: LocalDBComponent = component_class().set(**default_kwargs)
|
||||
component.build_vector_store()
|
||||
|
||||
# Switch to Retrieve mode
|
||||
component.set(mode="Retrieve", search_query="quick fox")
|
||||
results = component.search_documents()
|
||||
|
||||
assert len(results) == 3
|
||||
# Results should be diverse but relevant
|
||||
assert any("fox" in result.text.lower() for result in results)
|
||||
mock_mmr_search.assert_called_once_with(query="quick fox", k=3)
|
||||
|
||||
# Test with different settings
|
||||
component.set(number_of_results=2)
|
||||
mock_mmr_search.return_value = mock_docs[:2]
|
||||
diverse_results = component.search_documents()
|
||||
assert len(diverse_results) == 2
|
||||
|
||||
@patch("langchain_chroma.Chroma.similarity_search")
|
||||
@patch("langchain_chroma.Chroma.max_marginal_relevance_search")
|
||||
def test_search_with_different_types(
|
||||
self,
|
||||
mock_mmr_search,
|
||||
mock_similarity_search,
|
||||
component_class: type[LocalDBComponent],
|
||||
default_kwargs: dict[str, Any],
|
||||
) -> None:
|
||||
"""Test search with different search types."""
|
||||
test_data = [
|
||||
"The quick brown fox jumps over the lazy dog",
|
||||
"Python is a popular programming language",
|
||||
"Machine learning models process data",
|
||||
]
|
||||
default_kwargs["ingest_data"] = [Data(text=text) for text in test_data]
|
||||
default_kwargs["number_of_results"] = 2
|
||||
|
||||
# Mock the search methods to return documents
|
||||
from langchain_core.documents import Document
|
||||
|
||||
mock_similarity_docs = [
|
||||
Document(page_content="Python is a popular programming language"),
|
||||
Document(page_content="Machine learning models process data"),
|
||||
]
|
||||
mock_similarity_search.return_value = mock_similarity_docs
|
||||
|
||||
mock_mmr_docs = [
|
||||
Document(page_content="Python is a popular programming language"),
|
||||
Document(page_content="The quick brown fox jumps over the lazy dog"),
|
||||
]
|
||||
mock_mmr_search.return_value = mock_mmr_docs
|
||||
|
||||
component: LocalDBComponent = component_class().set(**default_kwargs)
|
||||
component.build_vector_store()
|
||||
|
||||
# Switch to Retrieve mode and test similarity search
|
||||
component.set(mode="Retrieve", search_type="Similarity", search_query="programming languages")
|
||||
similarity_results = component.search_documents()
|
||||
assert len(similarity_results) == 2
|
||||
assert any("python" in result.text.lower() for result in similarity_results)
|
||||
mock_similarity_search.assert_called_once_with(query="programming languages", k=2)
|
||||
|
||||
# Test MMR search
|
||||
component.set(search_type="MMR", search_query="programming languages")
|
||||
mmr_results = component.search_documents()
|
||||
assert len(mmr_results) == 2
|
||||
mock_mmr_search.assert_called_once_with(query="programming languages", k=2)
|
||||
|
||||
# Test with empty query
|
||||
component.set(search_query="")
|
||||
empty_results = component.search_documents()
|
||||
assert len(empty_results) == 0
|
||||
|
||||
@patch("langchain_chroma.Chroma.get")
|
||||
@patch("langchain_chroma.Chroma._collection")
|
||||
def test_duplicate_handling(
|
||||
self,
|
||||
mock_collection,
|
||||
mock_get,
|
||||
component_class: type[LocalDBComponent],
|
||||
default_kwargs: dict[str, Any],
|
||||
) -> None:
|
||||
"""Test handling of duplicate documents."""
|
||||
# Create test data with duplicates
|
||||
test_data = [
|
||||
Data(text_key="text", data={"text": "This is a test document"}),
|
||||
Data(text_key="text", data={"text": "This is a test document"}), # Duplicate with exact same data
|
||||
Data(text_key="text", data={"text": "This is another document"}),
|
||||
]
|
||||
default_kwargs["ingest_data"] = test_data
|
||||
default_kwargs["allow_duplicates"] = False
|
||||
default_kwargs["limit"] = 100 # Set a high enough limit to get all documents
|
||||
|
||||
# Mock the get method to return documents
|
||||
mock_get.return_value = {
|
||||
"documents": ["This is a test document", "This is a test document", "This is another document"],
|
||||
"metadatas": [{}, {}, {}],
|
||||
"ids": ["1", "2", "3"],
|
||||
}
|
||||
|
||||
# Mock collection count
|
||||
mock_collection.count.return_value = 3
|
||||
|
||||
component: LocalDBComponent = component_class().set(**default_kwargs)
|
||||
vector_store = component.build_vector_store()
|
||||
|
||||
# Get all documents
|
||||
results = vector_store.get(limit=100)
|
||||
documents = results["documents"]
|
||||
|
||||
# The documents are returned in a list structure
|
||||
assert len(documents) == 3 # All documents are added, even duplicates
|
||||
|
||||
# Count unique texts
|
||||
unique_texts = set(documents)
|
||||
assert len(unique_texts) == 2 # Should have 2 unique texts
|
||||
|
||||
# Test with allow_duplicates=True
|
||||
test_data = [
|
||||
Data(text_key="text", data={"text": "This is a test document"}),
|
||||
Data(text_key="text", data={"text": "This is a test document"}), # Duplicate
|
||||
]
|
||||
default_kwargs["ingest_data"] = test_data
|
||||
default_kwargs["allow_duplicates"] = True
|
||||
default_kwargs["collection_name"] = "test_collection_2" # Use a different collection name
|
||||
|
||||
# Mock for the second test
|
||||
mock_get.return_value = {
|
||||
"documents": ["This is a test document", "This is a test document"],
|
||||
"metadatas": [{}, {}],
|
||||
"ids": ["1", "2"],
|
||||
}
|
||||
mock_collection.count.return_value = 2
|
||||
|
||||
component = component_class().set(**default_kwargs)
|
||||
vector_store = component.build_vector_store()
|
||||
|
||||
# Get all documents
|
||||
results = vector_store.get(limit=100)
|
||||
documents = results["documents"]
|
||||
|
||||
# With allow_duplicates=True, we should have both documents
|
||||
assert len(documents) == 2
|
||||
assert all("test document" in doc for doc in documents)
|
||||
|
||||
# Verify that we have the expected number of documents
|
||||
assert vector_store._collection.count() == 2
|
||||
|
||||
def test_build_config_update(self, component_class: type[LocalDBComponent]) -> None:
|
||||
"""Test the update_build_config method."""
|
||||
component = component_class()
|
||||
|
||||
# Test mode=Ingest
|
||||
build_config = {
|
||||
"ingest_data": {"show": False},
|
||||
"collection_name": {"show": False},
|
||||
"persist": {"show": False},
|
||||
"persist_directory": {"show": False},
|
||||
"embedding": {"show": False},
|
||||
"allow_duplicates": {"show": False},
|
||||
"limit": {"show": False},
|
||||
"search_query": {"show": False},
|
||||
"search_type": {"show": False},
|
||||
"number_of_results": {"show": False},
|
||||
"existing_collections": {"show": False},
|
||||
}
|
||||
|
||||
updated_config = component.update_build_config(build_config, "Ingest", "mode")
|
||||
|
||||
assert updated_config["ingest_data"]["show"] is True
|
||||
assert updated_config["collection_name"]["show"] is True
|
||||
assert updated_config["persist"]["show"] is True
|
||||
assert updated_config["search_query"]["show"] is False
|
||||
|
||||
# Test mode=Retrieve
|
||||
updated_config = component.update_build_config(build_config, "Retrieve", "mode")
|
||||
|
||||
assert updated_config["search_query"]["show"] is True
|
||||
assert updated_config["search_type"]["show"] is True
|
||||
assert updated_config["number_of_results"]["show"] is True
|
||||
assert updated_config["existing_collections"]["show"] is True
|
||||
assert updated_config["collection_name"]["show"] is False
|
||||
|
||||
# Test persist=True/False
|
||||
build_config = {"persist_directory": {"show": False}}
|
||||
# Use keyword arguments to fix FBT003
|
||||
updated_config = component.update_build_config(build_config, field_value=True, field_name="persist")
|
||||
assert updated_config["persist_directory"]["show"] is True
|
||||
|
||||
updated_config = component.update_build_config(build_config, field_value=False, field_name="persist")
|
||||
assert updated_config["persist_directory"]["show"] is False
|
||||
|
||||
# Test existing_collections update
|
||||
# Fix the dict entry type issue
|
||||
build_config = {"collection_name": {"value": "old_name", "show": False}}
|
||||
updated_config = component.update_build_config(build_config, "new_collection", "existing_collections")
|
||||
assert updated_config["collection_name"]["value"] == "new_collection"
|
||||
|
||||
@patch("langflow.components.vectorstores.local_db.LocalDBComponent.list_existing_collections")
|
||||
def test_list_existing_collections(self, mock_list: MagicMock, component_class: type[LocalDBComponent]) -> None:
|
||||
"""Test the list_existing_collections method."""
|
||||
mock_list.return_value = ["collection1", "collection2", "collection3"]
|
||||
|
||||
component = component_class()
|
||||
collections = component.list_existing_collections()
|
||||
|
||||
assert collections == ["collection1", "collection2", "collection3"]
|
||||
mock_list.assert_called_once()
|
||||
Loading…
Add table
Add a link
Reference in a new issue